// This program is an attachment of the article 
// An infinite family of MUB-triplets in dimension 6
// written by P. Jaming, M. Matolcsi, P. M\'{o}ra,
// F. Sz\"{o}ll\H{o}si and M. Weiner.
//
// This is a C++ program. I compile it with gcc compiler under linux
// with
//
//  g++ -O3 ort.cpp -o ort
//
// command, no other files are required. There
// should be no problem compile it with a 32/64 bit compiler.
//
// This program is a part of a proof. In order to repeat this
// proof run this program with the following parameters:
// 
//  ./ort 19 7
//
// To finish this proof compile and run fab_ubv.cpp as well.
//
//
// Under Windows it can be compiled with Dev-Cpp (free compiler,
// http://www.bloodshed.net/devcpp.html). After opening this file
// you can compile and run this program with button F9. This program
// requires two command line parameters (otherwise it returns with
// error but you might not be able to read it because it closes too
// fast), you can set these parameters under menu Execute-> Parameters... 
// For example set the "Parameters to pass to your program" to: 19 7
// 
// 
// The aim of this program is to list certain vectors and
// save them in a txt file. In the notations of the paper, 
// these vectors constitute the set $ORT_{N'}$. 
// 
// The notations of the paper and this program correspond to each other 
// in the following way: 
//
// The parameter m_n2 here is N' in the paper.
//  
// In the paper vectors are given as columns, which is the traditional way in 
// mathematics. However, for simplicity in this program we deal with 
// vectors as being rows. 
//
// The second parameter of the program, 7, means that we check 7 generations 
// of descendants in the multiscale startegy. 
//
//
// In the following the vector
// ((0,q,w,e,r,t)) at level m_n2
// represents a set of vectors in dimension 6, which has first
// coordinate 0 and the 2nd, 3rd, 4th, 5th, 6th coordinates
// are in the intervals [q/m_n2,(q+1)/m_n2], [w/m_n2,(w+1)/m_n2],
// [e/m_n2,(e+1)/m_n2], [r/m_n2,(r+1)/m_n2],
// [t/m_n2,(t+1)/m_n2], respectively. (We use the notation double brackets
// ((0,q,w,e,r,t)) to note that some coordinates of the vector 
// represent intervals.)
// The variable m_n2
// is the first parameter of the program, in our proof
// it is 19. We save those vectors ((0,q,w,e,r,t)) to ort_19.txt
// (where 19 is the given parameter),
// for which there exists a vector in the set represented
// by ((0,q,w,e,r,t)), which is orthogonal to the vector
// (0,0,0,0,0,0) after applying t->e^{i*2*Pi*t} in both
// vectors' all coordinates.
// For simplicity we say that the vector 
// ((0,q,w,e,r,t)) is orthogonal to the vector (0,0,0,0,0,0)
// (although the first of these is a set, not a siingle vector). 
// We save all vectors ((0,q,w,e,r,t))
// which are orthogonal to the zero vector, but because of 
// the error estimates we might save some other vectors as well.
// The higher the second parameters of this program, 
// the less vectors are saved which are not orthogonal to the
// zero vector, and the more memory and time are used.
// 

#include <iostream>
#include <stdlib.h>
#include <stdio.h>
#include <math.h>
#include <vector>
#include <time.h>
#include <algorithm>
#include <stdlib.h>

using namespace std;


// For every calculation we use double as floating type.
// The precision of a double type variable is about 16 decimal digits,
// therefore all round-off error are less than EPSILON:=10^{-10}.
#define EPSILON 0.0000000001
#define PI2 6.28318530717958647692


// The function itostr converts an integer to string.
string   itostr(int i)
{
	char s[50];
	sprintf(s,"%d",i);
	return string(s);
};



// The array m_shift contains 5*32 values. 
// There is 2^5 ways to make a 0-1 list with length 5, m_shift consists of all of these.
// Namely:
// 0, 0, 0, 0, 0
// 0, 0, 0, 0, 1
// 0, 0, 0, 1, 0
// 0, 0, 0, 1, 1
// ...
// 1, 1, 1, 1, 1
int	m_shift [] = {0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 0, 1, 1, 0, 0, 1, 0, 0, 0, 0, 1, 0, 1, 0, 0, 1, 1, 0, 0, 0, 1, 1, 1, 0, 1, 0, 0, 0, 0, 1, 0, 0, 1, 0, 1, 0, 1, 0, 0, 1, 0, 1, 1, 0, 1, 1, 0, 0, 0, 1, 1, 0, 1, 0, 1, 1, 1, 0, 0, 1, 1, 1, 1, 1, 0, 0, 0, 0, 1, 0, 0, 0, 1, 1, 0, 0, 1, 0, 1, 0, 0, 1, 1, 1, 0, 1, 0, 0, 1, 0, 1, 0, 1, 1, 0, 1, 1, 0, 1, 0, 1, 1, 1, 1, 1, 0, 0, 0, 1, 1, 0, 0, 1, 1, 1, 0, 1, 0, 1, 1, 0, 1, 1, 1, 1, 1, 0, 0, 1, 1, 1, 0, 1, 1, 1, 1, 1, 0, 1, 1, 1, 1, 1};

// All the computations are done by the class Mub. 
// We will declare one instance of it.
class Mub{

	int		m_n2;

// m_max_n2 = m_n2 * 2^(q),
// where q is the second parameter and
// ^ is the power function.
	int		m_max_n2;

// We search for all those vectors at level m_n2 which 
// are orthogonal to (0,0,0,0,0,0). To do so, we use 
// the array m_row. The variable m_row[0] is always 0.
	int		m_row[6];

// The function scalar returns true if the current value of
// m_row at level n can be orthogonal to (0,0,0,0,0,0).
	bool		scalar(int n);

// The number of found cases.
	int		m_found_cases;

// For error estimate and scalar product we need to calculate a lot of values
// many times. To speed up the program we store these values in advance.

// In the function scalar we calculate a scalar product which
// uses error estimate.
// m_error[i] =  (double)5 * PI2 / (double) (2*i);
	double*		m_error;		

// For i=0,1,...,4*m_max_n2 we have
//  ssin[i] = sin(((double)i) * PI2 / (double)(2*m_max_n2));
//  ccos[i] = cos(((double)i) * PI2 / (double)(2*m_max_n2));
	double*		ssin;
	double*		ccos;

// The function start calls function iterate with n=m_n2.
	bool		iterate(int n);

public:
// The constructor.
	Mub(int n2,int max_n2);
	~Mub();

	void		start();
	int		found_cases() { return m_found_cases; }
};

Mub::Mub(int n2, int max_n2)
{
	m_found_cases = 0;

	m_n2 = n2;
	m_max_n2 = max_n2;
	ssin = new double[4*m_max_n2];
	ccos = new double[4*m_max_n2];
	for (int i = 0; i < 4*m_max_n2; i++)
	{
		ssin[i] = sin(((double)i) * PI2 / (double)(2*m_max_n2));
		ccos[i] = cos(((double)i) * PI2 / (double)(2*m_max_n2));
	}
	m_error = new double[m_max_n2+1];
	for (int i=n2; i <= m_max_n2; i*=2)
		m_error[i] =  (double)5 * PI2 / (double) (2*i);
};

Mub::~Mub()
{
	delete[] ssin;
	delete[] ccos;
	delete[] m_error;
};


void	Mub::start()
{
	
// We open a file for writing. In our case it is ort_19.txt.
	FILE*	out;
	string filename;
	filename = "ort_";
	filename = filename + itostr(m_n2) + ".txt";
	out = fopen(filename.c_str(),"w");	
	m_row[0] = 0;
	// We try out all possible choices for the array m_row, and
	// call the function iterate.
	for (m_row[1] = 0; m_row[1] < m_n2; m_row[1]++)
	{
		printf("%d/%d, found cases: %d\n",m_row[1],m_n2,m_found_cases);
		for (m_row[2] = 0; m_row[2] < m_n2; m_row[2]++)
			for (m_row[3] = 0; m_row[3] < m_n2; m_row[3]++)
				for (m_row[4] = 0; m_row[4] < m_n2; m_row[4]++)
					for (m_row[5] = 0; m_row[5] < m_n2; m_row[5]++)
					{
						if(iterate(m_n2))
						{
							// We found a vector, we write it out to the file.
							m_found_cases++;
							fprintf(out,"%d, %d, %d, %d, %d, %d\n",m_row[0],m_row[1],m_row[2],m_row[3],m_row[4],m_row[5]);
						}
					}
	}

	fclose(out);
}


bool	Mub::iterate(int n)
{
	// If ((0,m_row[1],m_row[2],m_row[3],m_row[4],m_row[5])) at level n
	// is not orthogonal to (0,0,0,0,0,0) then scalar(n) returns false.
	if (!scalar(n))
		return false;

	// If n < m_max_n2 then we run function iterate for all vectors
	// ((0,2*m_row[1]+a,2*m_row[2]+b,2*m_row[3]+c,2*m_row[4]+d,2*m_row[5]+e))
	// at level 2*n where a, b, c, d, e are 0 or 1.
	if (n < m_max_n2)
	{
		int r1=m_row[1]*2;
		int r2=m_row[2]*2;
		int r3=m_row[3]*2;
		int r4=m_row[4]*2;
		int r5=m_row[5]*2;

		bool ret = false;

		for (int i =0 ; i < 160; i+=5)
		{
			m_row[1]=r1+m_shift[i+0];
			m_row[2]=r2+m_shift[i+1];
			m_row[3]=r3+m_shift[i+2];
			m_row[4]=r4+m_shift[i+3];
			m_row[5]=r5+m_shift[i+4];
			
			if (iterate(n*2))
			{
				ret = true;
				break;
			}
		}

		m_row[1]=r1/2;
		m_row[2]=r2/2;
		m_row[3]=r3/2;
		m_row[4]=r4/2;
		m_row[5]=r5/2;
		return ret;
	}
	else
	{
		return true;
	}

}

bool	Mub::scalar(int n)
{
// We make the scalar product of the vector
// (0,(m_row[1]+0.5)/n,(m_row[2]+0.5)/n,(m_row[3]+0.5)/n,
// (m_row[4]+0.5)/n,(m_row[5]+0.5)/n) and (0,0,0,0,0,0)
// after applying t->e^{i*2*Pi*t} in all coordinates.
	double x = 1;
	double y = 0;
	int mult = m_max_n2/n;
	int mult_for_m_n2 = m_max_n2/m_n2;
	double r;
	for ( int i = 1; i < 6; i++)
	{
		x += ccos[ (2*m_row[i]+1)*(mult) ];
		y += ssin[ (2*m_row[i]+1)*(mult) ];
	}
// The vector ((0,m_row[1],m_row[2],m_row[3],m_row[4],m_row[5]))
// at level n represents a set. We only calculated the scalar product
// of a vector of this set. The error we made this way is stored in
// m_error[n].
	if ( (r=sqrt(x*x+y*y)) > m_error[n] + EPSILON )
	{
		return false;
	}


// We make a more sophisticated error estimate. If n > 15 then
// m_error[n] < 1. The "error" (the distance from 0+i*0) is stored
// in variable r. We check whether this error r can be realized
// between x+i*y and the scalar product of an arbitrary vector v 
// in the set represented by ((0,m_row[1],m_row[2],m_row[3],m_row[4],m_row[5])) 
// at level n and (0,0,0,0,0,0).
// We do not have to check each vectors v, because the m_error[n]<1 thus
// the maximum distance between x+i*y and the scalar product of
// the vector v and (0,0,0,0,0,0) is achieved when the vector v
// is extremal. Namely, n*v has only integer coordinates.
	double xx,yy;

	if (n > 15)
	{
		for (int i = 0; i < 160; i+=5)
		{	
			xx=1.0+ccos[ 2*(m_row[1]+m_shift[i])*mult ]+
				ccos[ 2*(m_row[2]+m_shift[i+1])*mult ]+
				ccos[ 2*(m_row[3]+m_shift[i+2])*mult ]+
				ccos[ 2*(m_row[4]+m_shift[i+3])*mult ]+
				ccos[ 2*(m_row[5]+m_shift[i+4])*mult ] - x;
			yy=ssin[ 2*(m_row[1]+m_shift[i])*mult ]+
				ssin[ 2*(m_row[2]+m_shift[i+1])*mult ]+
				ssin[ 2*(m_row[3]+m_shift[i+2])*mult ]+
				ssin[ 2*(m_row[4]+m_shift[i+3])*mult ]+
				ssin[ 2*(m_row[5]+m_shift[i+4])*mult ] - y;
			if (r-sqrt(xx*xx+yy*yy) < EPSILON)
			{
				return true;
			}
		}
		return false;
	}
	else
		return true;
}


int	main(int argc, const char* argv[]){

	if (argc != 3)
	{
		printf("Error, usage: ort n how_many_descendants\nwhere how_many_descendants>0 \n");
		return 1;
	}

	int max_n2 = atoi(argv[1]);
	int q = atoi(argv[2]);
	for (int i = 0; i < q; i++)
		max_n2 *= 2;

	Mub mub(atoi(argv[1]),max_n2);

	printf("START\n");


	mub.start();
	printf("\nFound cases: %d\n",mub.found_cases());

	return 0;
}
