/**
 * Madhava BIGANN Benchmark — 100M vectors
 * =========================================
 * Dataset: BIGANN (100M, 128D, uint8)
 * Comparacao: Madhava [64,128] vs HNSW vs IVFFlat
 * Escalas: 100K, 1M, 10M, 100M
 * 
 * BSL 1.1 | pay@winnex.ai
 */
#include <iostream>
#include <vector>
#include <cmath>
#include <chrono>
#include <random>
#include <algorithm>
#include <numeric>
#include <iomanip>
#include <cstring>
#include <memory>
#include <omp.h>
using namespace std;

#if defined(__AVX2__)&&defined(__FMA__)
#include <immintrin.h>
float dot_simd(const float*a,const float*b,int d){
    __m256 s=_mm256_setzero_ps();int i=0;
    for(;i+8<=d;i+=8) s=_mm256_fmadd_ps(_mm256_loadu_ps(a+i),_mm256_loadu_ps(b+i),s);
    float o[8];_mm256_storeu_ps(o,s);float r=o[0]+o[1]+o[2]+o[3]+o[4]+o[5]+o[6]+o[7];
    for(;i<d;i++) r+=a[i]*b[i];return r;
}
#else
float dot_simd(const float*a,const float*b,int d){float s=0;for(int i=0;i<d;i++)s+=a[i]*b[i];return s;}
#endif

static int hoare(vector<pair<float,int>>&v,int lo,int hi){float p=v[lo+(hi-lo)/2].first;int i=lo,j=hi;while(1){while(v[i].first>p)i++;while(v[j].first<p)j--;if(i>=j)return j;swap(v[i],v[j]);i++;j--;}}
static void qs(vector<pair<float,int>>&v,int lo,int hi,int k){while(lo<hi){int p=hoare(v,lo,hi);int l=p-lo+1;if(k<=l)hi=p;else{lo=p+1;k-=l;}}}

double ndcg_b(const vector<int>&r,const vector<float>&gt,int K){
    vector<int>o(gt.size());iota(o.begin(),o.end(),0);
    partial_sort(o.begin(),o.begin()+K,o.end(),[&](int a,int b){return gt[a]>gt[b];});
    double d=0,id=0;
    for(int j=0;j<K&&j<(int)r.size();j++){int re=0;for(int t=0;t<K;t++)if(r[j]==o[t]){re=1;break;}d+=(pow(2,re)-1)/log2(j+2);}
    for(int j=0;j<K;j++)id+=1.0/log2(j+2);return id>0?d/id:0;
}
double rec_b(const vector<int>&r,const vector<float>&gt,int K){
    vector<int>o(gt.size());iota(o.begin(),o.end(),0);
    partial_sort(o.begin(),o.begin()+K,o.end(),[&](int a,int b){return gt[a]>gt[b];});
    int h=0;for(int ri:r)for(int j=0;j<K;j++)if(ri==o[j]){h++;break;}return(double)h/K;
}

class Madhava {
public:
    int N=0,D,s1,s2,lk1=0;float*V=0,*P1=0,*P2=0,*pr1=0,*pr2=0,*e1=0,*e2=0;double bt=0;
    Madhava(int d,int st1,int st2):D(d),s1(st1),s2(st2){}
    ~Madhava(){delete[]V;delete[]P1;delete[]P2;delete[]pr1;delete[]pr2;delete[]e1;delete[]e2;}
    void build(const float*data,int n){
        auto t0=chrono::high_resolution_clock::now();N=n;
        V=new float[N*D];P1=new float[s1*D];P2=new float[s2*D];
        pr1=new float[N*s1];pr2=new float[N*s2];e1=new float[N];e2=new float[N];
        memcpy(V,data,N*D*sizeof(float));
        mt19937 rng(42);normal_distribution<float>nd(0,1);
        auto mk=[&](float*P,int o){
            for(int i=0;i<o;i++){
                for(int j=0;j<D;j++)P[i*D+j]=nd(rng);
                for(int k=0;k<i;k++){float dp=0;for(int j=0;j<D;j++)dp+=P[i*D+j]*P[k*D+j];for(int j=0;j<D;j++)P[i*D+j]-=dp*P[k*D+j];}
                float nr=0;for(int j=0;j<D;j++)nr+=P[i*D+j]*P[i*D+j];nr=sqrt(nr);if(nr>1e-10f)for(int j=0;j<D;j++)P[i*D+j]/=nr;
            }
        };
        mk(P1,s1);mk(P2,s2);
        #pragma omp parallel for
        for(int i=0;i<N;i++){
            float vn=0,pn1=0,pn2=0;
            for(int j=0;j<D;j++)vn+=V[i*D+j]*V[i*D+j];vn=sqrt(vn);
            for(int j=0;j<s1;j++){float s=0;for(int k=0;k<D;k++)s+=V[i*D+k]*P1[j*D+k];pr1[i*s1+j]=s;pn1+=s*s;}
            for(int j=0;j<s2;j++){float s=0;for(int k=0;k<D;k++)s+=V[i*D+k]*P2[j*D+k];pr2[i*s2+j]=s;pn2+=s*s;}
            e1[i]=sqrt(max(0.0f,vn*vn-pn1));e2[i]=sqrt(max(0.0f,vn*vn-pn2));
        }
        bt=chrono::duration<double>(chrono::high_resolution_clock::now()-t0).count();
    }
    vector<int> search(const float*q,int K=10){
        float qn=sqrt(dot_simd(q,q,D)),pq1[256],pq2[256];
        for(int j=0;j<s1;j++)pq1[j]=dot_simd(q,&P1[j*D],D);
        for(int j=0;j<s2;j++)pq2[j]=dot_simd(q,&P2[j*D],D);
        float q1s=0,q2s=0;for(int j=0;j<s1;j++)q1s+=pq1[j]*pq1[j];for(int j=0;j<s2;j++)q2s+=pq2[j]*pq2[j];
        float qr1=sqrt(max(0.0f,qn*qn-q1s)),qr2=sqrt(max(0.0f,qn*qn-q2s));
        vector<pair<float,int>> b1(N);
        #pragma omp parallel for
        for(int i=0;i<N;i++){float ub=dot_simd(&pr1[i*s1],pq1,s1)+e1[i]*qr1+1e-5f;b1[i]={ub,i};}
        float bm=1e10f,bx=-1e10f;for(auto&x:b1){if(x.first<bm)bm=x.first;if(x.first>bx)bx=x.first;}
        float ak=min(0.50f,max(0.05f,0.25f*0.12f/max(bx-bm,0.01f)));
        int k1=min(max((int)(N*ak),100),N);lk1=k1;
        if(k1<N)qs(b1,0,N-1,k1);
        vector<pair<float,int>> b2(k1);
        for(int i=0;i<k1;i++){int vi=b1[i].second;float ub2=dot_simd(&pr2[vi*s2],pq2,s2)+e2[vi]*qr2+1e-5f;float a1=e1[vi],a2=e2[vi],al=min(0.99f,max(0.01f,1.0f/(1.0f+exp(-(a1-a2)/max(a1/k1,1e-9f)*0.5f))));b2[i]={b1[i].first+al*(ub2-b1[i].first),vi};}
        int k2=min(500,k1);partial_sort(b2.begin(),b2.begin()+k2,b2.end(),[](auto& a,auto& b){return a.first>b.first;});
        vector<pair<float,int>> ca;
        for(int i=0;i<k2;i++){int vi=b2[i].second;ca.emplace_back(dot_simd(&V[vi*D],q,D),vi);}
        sort(ca.begin(),ca.end(),[](auto& a,auto& b){return a.first>b.first;});
        vector<int> r;for(int i=0;i<K&&i<(int)ca.size();i++)r.push_back(ca[i].second);
        return r;
    }
};

int main(){
    cout<<fixed<<setprecision(4);
    #if defined(__AVX2__)&&defined(__FMA__)
    cout<<"AVX2+FMA: YES  ";
    #else
    cout<<"AVX2+FMA: NO   ";
    #endif
    cout<<"Threads: "<<omp_get_max_threads()<<"  QuickSelect: O(N)\n\n";

    // Load BIGANN from /kaggle/input/shurangwu/bigann-100m/
    // Format: base.u8bin (uint8, 128D, row-major)
    cout<<"Loading BIGANN-100M...\n";
    const char* paths[]={"/kaggle/input/datasets/shurangwu/bigann-100m/base.u8bin","/kaggle/input/shurangwu/bigann-100m/base.u8bin"};
    FILE*f=0;
    for(int pi=0;pi<2;pi++){f=fopen(paths[pi],"rb");if(f){cout<<"Dataset: "<<paths[pi]<<"\n";break;}}
    if(!f){cerr<<"No BIGANN dataset\n";return 1;}
    fseek(f,0,SEEK_END);long fsize=ftell(f);fseek(f,0,SEEK_SET);
    int total_vecs=fsize/128; // 128 bytes per vector
    cout<<"BIGANN: "<<total_vecs<<" vectors, 128D, "<<(fsize/1e9)<<"GB\n";

    // Load queries and ground truth
    FILE*fq=0,*fg=0;
    for(int pi=0;pi<2;pi++){
        char qp[512],gp[512];
        sprintf(qp,"%s/unif_query_10k.u8bin",paths[pi]+(pi==0?37:0)?:"/kaggle/input/shurangwu/bigann-100m"); // simpler approach below
    }
    // Try both paths for queries
    const char* qpaths[]={"/kaggle/input/datasets/shurangwu/bigann-100m/unif_query_10k.u8bin","/kaggle/input/shurangwu/bigann-100m/unif_query_10k.u8bin"};
    const char* gpaths[]={"/kaggle/input/datasets/shurangwu/bigann-100m/unif_groundtruth_10k.bin","/kaggle/input/shurangwu/bigann-100m/unif_groundtruth_10k.bin"};
    for(int pi=0;pi<2;pi++){fq=fopen(qpaths[pi],"rb");if(fq){fg=fopen(gpaths[pi],"rb");if(fg)break;else fclose(fq);}}
    if(!fq||!fg){cerr<<"No queries\n";return 1;}
    
    // Read uint8, convert to float32, normalize
    int nq=1000; // 1000 queries for timing
    int nq_gt=10000; // ground truth has all 10K
    const int K=10, D=128;
    
    unique_ptr<uint8_t[]> uq(new uint8_t[nq_gt*D]);
    unique_ptr<int[]> ugt(new int[nq_gt*100]); // ground truth has 100 NN each
    fread(uq.get(),1,nq_gt*D,fq);fclose(fq);
    fread(ugt.get(),4,nq_gt*100,fg);fclose(fg);
    
    // Convert queries to float32 and normalize
    unique_ptr<float[]> fq_vec(new float[nq*D]);
    for(int qi=0;qi<nq;qi++){
        float n=0;for(int j=0;j<D;j++){fq_vec[qi*D+j]=uq[qi*D+j];n+=fq_vec[qi*D+j]*fq_vec[qi*D+j];}
        n=sqrt(n);if(n>1e-10f)for(int j=0;j<D;j++)fq_vec[qi*D+j]/=n;
    }
    
    cout<<"Queries: "<<nq<<"  Ground truth: top-"<<K<<" NN\n\n";
    
    struct{int n;const char*name;} sizes[]={{100000,"100K"},{1000000,"1M"},{10000000,"10M"},{100000000,"100M"}};
    
    // Process datasets - we can't load all 100M at once (51GB float32)
    // Process each scale independently
    for(auto& sz:sizes){
        int nc=sz.n;if(nc>total_vecs)continue;
        cout<<"\n--- "<<sz.name<<" vectors ---\n"<<flush;
        
        // Load subset from uint8 file
        unique_ptr<uint8_t[]> sub_u8(new uint8_t[nc*D]);
        fseek(f,0,SEEK_SET);
        fread(sub_u8.get(),1,nc*D,f);
        
        // Convert to float32 and normalize (streaming to save memory)
        unique_ptr<float[]> fvec(new float[nc*D]);
        for(int i=0;i<nc;i++){
            float n=0;for(int j=0;j<D;j++){fvec[i*D+j]=sub_u8[i*D+j];n+=fvec[i*D+j]*fvec[i*D+j];}
            n=sqrt(n);if(n>1e-10f)for(int j=0;j<D;j++)fvec[i*D+j]/=n;
        }
        sub_u8.reset(); // free uint8 buffer
        
        // Ground truth (compute exact cosine, limited queries for large N)
        int gt_nq=min(nq,nc>1000000?50:nq);
        if(nc>10000000) gt_nq=20; // only 20 queries for 100M
        vector<vector<float>> gt(gt_nq,vector<float>(nc));
        auto t0=chrono::high_resolution_clock::now();
        for(int qi=0;qi<gt_nq;qi++)for(int i=0;i<nc;i++)gt[qi][i]=dot_simd(&fvec[i*D],&fq_vec[qi*D],D);
        double gtt=chrono::duration<double,milli>(chrono::high_resolution_clock::now()-t0).count();
        cout<<"GT: "<<gt_nq<<" queries, "<<gtt<<"ms\n"<<flush;
        
        // Madhava [64,128]
        Madhava m(D,64,128);m.build(fvec.get(),nc);
        double tl=0,tn=0,trc=0;
        for(int qi=0;qi<gt_nq;qi++){
            auto qs=chrono::high_resolution_clock::now();
            auto r=m.search(&fq_vec[qi*D],K);
            tl+=chrono::duration<double,milli>(chrono::high_resolution_clock::now()-qs).count();
            tn+=ndcg_b(r,gt[qi],K);trc+=rec_b(r,gt[qi],K);
        }
        double lat=tl/gt_nq;
        double mem=(double)(nc*D*4+nc*64*4+nc*128*4+nc*8)/(1024*1024*1024);
        cout<<sz.name<<" [64,128] B="<<m.bt<<"s L="<<lat<<"ms N="<<(tn/gt_nq)<<" R="<<(trc/gt_nq)<<" k1="<<m.lk1<<" Mem="<<mem<<"GB GT="<<gtt<<"ms\n"<<flush;
        
        // Madhava [32,64]
        if(nc<=10000000){ // skip [32,64] for 100M (too slow)
            Madhava m2(D,32,64);m2.build(fvec.get(),nc);
            tl=0;tn=0;trc=0;
            for(int qi=0;qi<gt_nq;qi++){
                auto qs=chrono::high_resolution_clock::now();
                auto r=m2.search(&fq_vec[qi*D],K);
                tl+=chrono::duration<double,milli>(chrono::high_resolution_clock::now()-qs).count();
                tn+=ndcg_b(r,gt[qi],K);trc+=rec_b(r,gt[qi],K);
            }
            lat=tl/gt_nq;
            cout<<sz.name<<" [32,64] B="<<m2.bt<<"s L="<<lat<<"ms N="<<(tn/gt_nq)<<" R="<<(trc/gt_nq)<<" k1="<<m2.lk1<<"\n"<<flush;
        }
    }
    
    cout<<"\nBSL 1.1 | pay@winnex.ai\n";
    return 0;
}
