# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# This Python 3 environment comes with many helpful analytics libraries installed
# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python
# For example, here's several helpful packages to load

import numpy as np # linear algebra
import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)

import matplotlib.pyplot as plt

# Input data files are available in the read-only "../input/" directory
# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory

import os
#for dirname, _, filenames in os.walk('/kaggle/input'):
    #for dn in dirname:
#    print(dirname);
#        print(os.path.join(dirname, filename))


r1 = 6378137.0;

fl = 1 / 298.257;

r2 = r1 - r1*fl;

d1_err = 2*np.pi*r1 / 360; ## Err E/W at equator // ~ N/S


r1_c = r1 + 7797; #7398
r2_c = r2 - 13564; #13962

##---------

WGS84_SEMI_MAJOR_AXIS = 6378137.0
WGS84_SEMI_MINOR_AXIS = 6356752.314245
WGS84_SQUARED_FIRST_ECCENTRICITY  = 6.69437999013e-3
WGS84_SQUARED_SECOND_ECCENTRICITY = 6.73949674226e-3


def ECEF_to_BLH(ecef):
    a = WGS84_SEMI_MAJOR_AXIS
    b = WGS84_SEMI_MINOR_AXIS
    e2  = WGS84_SQUARED_FIRST_ECCENTRICITY
    e2_ = WGS84_SQUARED_SECOND_ECCENTRICITY
    #x = ecef.x
    #y = ecef.y
    #z = ecef.z
    x = ecef[0]
    y = ecef[1]
    z = ecef[2]
    
    r = np.sqrt(x**2 + y**2)
    t = np.arctan2(z * (a/b), r)
    B = np.arctan2(z + (e2_*b)*np.sin(t)**3, r - (e2*a)*np.cos(t)**3)
    L = np.arctan2(y, x)
    n = a / np.sqrt(1 - e2*np.sin(B)**2)
    H = (r / np.cos(B)) - n
    #return BLH(lat=B, lng=L, hgt=H)
    
    return [B, L, H];

###----------------

from astropy import units as u
from astropy.coordinates import EarthLocation as loca

def ecef2ll(x,y,z):
    l = loca.from_geocentric(x,y,z,u.m).to_geodetic();
    return [l.lat.value, l.lon.value];



def xlin(a,b,x):
    return (1-x)*a + x*b;

def xquad(a,b,c,x):  ##smooth b?
    return xlin(xlin(a,b,x), xlin(b,c,x), x);

def xquad2(a,b,c,x):  ##
    b2 = xlin((a+c)/2,b,2);
    return xlin(xlin(a,b2,x), xlin(b2,c,x), x);


def xcubic(a,b,c,d,x):
    return xlin(xquad(a,b,c,x), xquad(b,c,d,x), x);

def xcubic2(a,b,c,d,x):
    return xlin(xquad2(a,b,c,x), xquad2(b,c,d,x), x);



dir1 = '/kaggle/input/smartphone-decimeter-2023/sdc2023/test'

l1 = [];

dirsub = '/kaggle/input/smartphone-decimeter-2023/sdc2023/sample_submission.csv';

sub1 = pd.read_csv(dirsub);

sub2 = pd.read_csv(dirsub);

print("submission", sub1.shape);

latlong_mat = np.zeros((sub1.shape[0], 2));

for dirpath, dirnames, filenames in os.walk(dir1):
    #for dn in dirname:
    #print(dirnames);
    l1 = dirnames;
    break;
    #os.walk()

    
print(len(l1));
    
l2 = ["pixel4/", "pixel5/", "mi8/", "pixel6pro"]
l3 = ["ground_truth.csv", "device_gnss.csv"];

fig = plt.figure()
ax = fig.add_subplot(projection='3d')


fig2 = plt.figure()
ax2 = fig2.add_subplot(projection='3d')


fig3 = plt.figure()
ax3 = fig3.add_subplot(projection='3d')

fig4 = plt.figure()
ax4 = fig4.add_subplot();

n_case = len(l1);

for data_id in range(n_case): #=1;
    
    
    #data_id=6; ##
    
    print("iter", data_id);

#train1 = dir1 + "/" + l1[data_id] + "/" + l2[1] + l3[0];

    xa = os.walk(dir1 + "/" + l1[data_id]);

    xa = [x for x in xa];

    xb = xa[0][1][0];
    
    n_ph = len(xa[0][1]);
    
    tx1 = l1[data_id] +'/'+xb
    
    sub1_v = sub1['tripId'] == tx1;
    
#print(xb);

    train1 = dir1 + "/" + l1[data_id] + "/" + xb + "/" + l3[0];

    train2 = dir1 + "/" + l1[data_id] + "/" + xb + "/" + l3[1];

    print("skip reading", train1);

    #t1 = pd.read_csv(train1);

    print("reading", train2);

    t2 = pd.read_csv(train2);

    t1 = [];
    
    print(t2.shape);

    #print(t1.head());

    #print(t2.head());
    
    #if t2["WlsPositionZEcefMeters"][0] > 3700000:

        #ax.scatter(t2["WlsPositionXEcefMeters"][::250],t2["WlsPositionYEcefMeters"][::250],t2["WlsPositionZEcefMeters"][::250]);

    """
    ax.scatter(t2["SvPositionXEcefMeters"][::125],t2["SvPositionYEcefMeters"][::125],t2["SvPositionZEcefMeters"][::125]);
    
    ax2.scatter(t2["SvVelocityXEcefMetersPerSecond"][::125],t2["SvVelocityYEcefMetersPerSecond"][::125],t2["SvVelocityZEcefMetersPerSecond"][::125]);
    
    ax3.scatter(t2["WlsPositionXEcefMeters"][::125],t2["WlsPositionYEcefMeters"][::125],t2["WlsPositionZEcefMeters"][::125]);
    
    ax4.scatter( t1["LongitudeDegrees"][::10], t1["LatitudeDegrees"][::10],);
    
    """    ;

#print(t1.columns);
#print(t2.columns);

    n2 = t2.shape[0];
    
    #n = t1.shape[0];
    
    ts = pd.unique(t2['utcTimeMillis']);
    
    n = ts.shape[0];
    
    train_err = np.zeros(n);
    
    n3 = np.sum(sub1_v);
    
    n3_targ = sub1['UnixTimeMillis'][sub1_v];
    
    ts2 = pd.unique(n3_targ);
    
    ts3 = n3_targ.shape[0];
    
    print("sub1_v size: ", ts3);
    
    print(ts.shape, ts2.shape, "assert?");
    
    xsub1_Lat = sub1['LatitudeDegrees'][sub1_v];
    
    xsub1_Long = sub1['LongitudeDegrees'][sub1_v];
    
    zsub1_Lat = np.zeros(xsub1_Lat.shape);
    
    zsub1_Long = np.zeros(xsub1_Lat.shape);
    
    print(xsub1_Lat.shape, n3, "match?");
    
    last_id=0;

    id_m4 =0;
    id_m3 =0;
    id_m2 =0;
    id_m1=0;
    
    pr_plots = np.zeros((n3, 9));
    
    spd_est = [];

    for test_inp in range(n3): #= 10;
    
        
        #test_inp=50;
        
        t_inp2 = last_id; #test_inp*35; #~
        
        if t_inp2>n2:
            t_inp2 = n2-50;
        
        
        while(ts[test_inp] > t2['utcTimeMillis'][t_inp2]): ##t1['UnixTimeMillis'][test_inp]
            t_inp2 += 5 * (t_inp2<=(n2-10));
            if t_inp2>=(n2-15):
                print("B?")
                break;

        while(ts[test_inp] < t2['utcTimeMillis'][t_inp2]): #while(t1['UnixTimeMillis'][test_inp]< t2['utcTimeMillis'][t_inp2]):
            t_inp2 -= 5;
            if t_inp2<10:
                break;
            
        last_id = t_inp2;
        
        if (test_inp > 0):
            
            id_m4 = id_m3;
            id_m3 = id_m2;
            id_m2 = id_m1;        
            id_m1 = np.array(targ1);
        
        targ1 = [t2["WlsPositionXEcefMeters"][t_inp2],t2["WlsPositionYEcefMeters"][t_inp2],t2["WlsPositionZEcefMeters"][t_inp2]];
        
        base2 = ecef2ll(targ1[0], targ1[1], targ1[2]);
        
        #targ_E1 = (targ1[2]**2 / r2_c**2) + (np.hypot(targ1[0], targ1[1])**2 / r1_c**2);
        
        targ_E1 = (targ1[2]**2 / r2**2) + (np.hypot(targ1[0], targ1[1])**2 / r1**2);
        
        targ1 = np.array(targ1) * (1/np.sqrt(targ_E1)); #((1.0+(1.0 / targ_E1))/2);
        
        targ_E2 = (targ1[2]**2 / r2**2) + (np.hypot(targ1[0], targ1[1])**2 / r1**2);
        
        #print("Elliptic vals", targ_E1, targ_E2, 1-targ_E1, 1-targ_E2, r1*(1-targ_E1), r1*(1-targ_E2));
        
        #if test_inp>1 and np.isnan(y1)==False:
        #    oldy1 = y1;
        #    y1 = np.hypot(np.hypot(targ1[0], targ1[1]), targ1[2]); ## vector len
        #    targ1 = np.array(targ1) * (((oldy1+y1)/2) / y1);
        
        
        
        if (test_inp>4):
            
            targ2 = 2*id_m1 - id_m2;
            
            targ_m1 = (targ1 + id_m2)/2;
            
            targ_m2 = (targ1 + id_m4)/2;
            
            #targ_m2 = (id_m1 + id_m3)/2;
            
            xtarg1 = 2*targ_m1 - targ_m2;
            
            xtarg1_b = xquad(id_m3, id_m2, id_m1, 1.5)
            
            xtarg1_c = xquad(id_m3, id_m1, xlin(id_m1, np.array(targ1),2), 0.75)
            
            xtarg1_d = xquad2(id_m3, id_m2, id_m1, 1.5)
            
            xtarg1_e = xquad2(id_m3, id_m1, xlin(id_m1, np.array(targ1),2), 0.75)
            
            xtarg1_f = xlin(xtarg1_c, xtarg1_e, 0.33);
            
            xtarg1_g = xcubic2(id_m4, id_m3, id_m2, id_m1, 1+(1/3));
            
            xtarg1_h = xcubic(id_m4, id_m3, id_m2, id_m1, 1+(1/3));
            
            if False:
            
                print("Projection 1 Err", np.sum(np.abs(xtarg1 - targ1)));
            
                print("Projection 2 Err", np.sum(np.abs(targ_m1 - id_m1)));
            
                print("Projection 3 Err", np.sum(np.abs(targ_m2 - id_m2)));
                
                print("Projection 4 Err", np.sum(np.abs(targ1 - xtarg1_b)));
                
                print("Projection 4 Err", np.sum(np.abs(targ1 - xtarg1_c)));
                
                print("Projection 5 Err", np.sum(np.abs(targ1 - xtarg1_d)));
                
                print("Projection 5 Err", np.sum(np.abs(targ1 - xtarg1_e)));
                
                print("static Err", np.sum(np.abs(targ1 - id_m1)));
            
                stat_err = np.sum(np.abs(targ1 - id_m1));
                
                print("Projection cubic2 Err", np.sum(np.abs(targ1 - xtarg1_g)), np.sum(np.abs(targ1 - xtarg1_g))/stat_err);
                
                print("Projection cubic1 Err", np.sum(np.abs(targ1 - xtarg1_h)), np.sum(np.abs(targ1 - xtarg1_h)) / stat_err);
                
                print("static / 100km/h", stat_err / (100000 / 3600));
            
            stat_err = np.sum(np.abs(targ1 - id_m1));
            
            spd_est.append(stat_err / (100000 / 3600));
                
            weight1 = 3 / np.maximum(np.sum(np.abs(targ1 - xtarg1_e)), 0.3); ## 50% @ 3m err? // max weight ~ 90%
            
            #pr_plots[test_inp, :] = [np.sum(np.abs(xtarg1 - targ1)), np.sum(np.abs(targ_m1 - id_m1)), np.sum(np.abs(targ_m2 - id_m2)), np.sum(np.abs(targ1 - xtarg1_b)), 
            #                        np.sum(np.abs(targ1 - xtarg1_c)), np.sum(np.abs(targ1 - xtarg1_d)), np.sum(np.abs(targ1 - xtarg1_e)), np.sum(np.abs(targ1 - xtarg1_f))];
            
            
            pr_plots[test_inp, 0:3] = targ1 - xtarg1_g;
            
            pr_plots[test_inp, 3:6] = targ1 - xtarg1_h;
            
            pr_plots[test_inp, 6:9] = targ1 - id_m1;
            
            #if np.sum(np.isnan(xtarg1))==0:
            #    targ1 = (targ1 + xtarg1_e*weight1) / (1 + weight1);

        base_x = ECEF_to_BLH(targ1);
        
        base_z1 = base_x[0];
        
        base_x1 = base_x[1];
        
        x1 = np.arctan2(targ1[1], targ1[0]) / np.pi*180; ##longitude

        y1 = np.hypot(np.hypot(targ1[0], targ1[1]), targ1[2]); ## vector len

        xy = np.hypot(targ1[0], targ1[1]);

        z1 = np.arctan2( targ1[2], xy) / np.pi*180; ##Altit xy / z ## (|> targ2 xy y1)
    
        z2 = np.arcsin(targ1[2] / r2) / np.pi*180; ## altit z (|> z r2)
        
        ## z2 + (z2-z1)
    
        z3 = np.arccos(xy / r2) / np.pi*180; ## xy min radius (|> xy r2)
        
        ## ~~ z1 + (z1 - z3)
    
        z4 = np.arccos(xy / r1) / np.pi*180; ## xy max radius (|> xy r1)
        
        ## ~~ z4 + (z4-z3)
        
        z6 = np.arccos(xy / y1) / np.pi*180; ## (|> targ2 xy y1)
        
        z62 = np.arcsin(targ1[2] / y1) / np.pi*180; ## (|> targ2 xy y1)
        
        z63 = np.arcsin(targ1[2] / r1) / np.pi*180; ## (|> z r1)
    
        z5 = (z2+z4)/2;
    
        z5 += z5 - z1;
    
        z5b = z2 + (z2-z1);
        
        z5 = (z5+z5b) / 2;
        
        w1 = np.arccos(xy / r1_c) / np.pi*180;
        
        w2 = np.arcsin(targ1[2] / r2_c) / np.pi*180;
        
        w3 = (w1 + 2*w2) / 3;
        
        #errX = t1["LatitudeDegrees"][test_inp];
        
        #adj1 = errX - z5;
        
        #adj2 = z2 - z1;
        
        if False:
            print("p2: ", adj1/adj2);
    
            print(z2, "est2?", z1, "est1?", z3, "est3", z4, "est4", z5, "e5", z5b, "est5b");
        
            print(z6, "estz6", z62, "est z62", z63, "z63");

            if test_inp==0:
                print(targ1, x1, y1, y1/r2, y1/r1);

                print(targ1[2],xy, z1);
    
            print(x1, z1, "[", z2,z3,z4,"]", z5);
    
            #print("Real targ", [t1["LatitudeDegrees"][test_inp], t1["LongitudeDegrees"][test_inp]]);
        
            #print(t1['UnixTimeMillis'][test_inp], t2['utcTimeMillis'][t_inp2], (t1['UnixTimeMillis'][test_inp]- t2['utcTimeMillis'][t_inp2])/1000, "time checks");
    
    
        #e1 = (t1["LongitudeDegrees"][test_inp] - x1) * d1_err;
    
        #e2 = (t1["LatitudeDegrees"][test_inp] - z5) * d1_err;
    
        #print("\n est dist errs", e2, e1, np.hypot(e1, e2), "\n\n");
        
        #train_err[test_inp] = np.hypot(e1, e2);
        
        #zsub1_Lat[test_inp] = w3; #sub1['LatitudeDegrees'][sub1_v];
    
        #zsub1_Long[test_inp] = x1; #sub1['LongitudeDegrees'][sub1_v];
        
        ## baseline test
        
        #zsub1_Lat[test_inp] = np.rad2deg(base_z1); #sub1['LatitudeDegrees'][sub1_v];
    
        zsub1_Long[test_inp] = np.rad2deg(base_x1); #sub1['LongitudeDegrees'][sub1_v];
        
        zsub1_Lat[test_inp] = w3;
        
        
        zsub1_Long[test_inp] = base2[1];
        
        zsub1_Lat[test_inp] = base2[0];
        

        
        if test_inp%1000 ==0:
            print(test_inp, xsub1_Lat.shape);
        
        
    #sub1['LatitudeDegrees'][sub1_v] = xsub1_Lat;
    
    print(xsub1_Lat.shape, zsub1_Lat.shape, np.sum(sub1_v));
    
    latlong_mat[sub1_v, 0] = zsub1_Lat;
    
    latlong_mat[sub1_v, 1] = zsub1_Long;
    
    #sub1['LongitudeDegrees'][sub1_v] = xsub1_Long;
    
    if data_id==0:
        
        plt.figure();
        plt.plot(pr_plots[:200,6]);
        plt.plot(pr_plots[:200,4]);
        plt.plot(pr_plots[:200,7]);
        plt.savefig("projection_err.jpg");
        
    ax.scatter(pr_plots[::100, 1], pr_plots[::100, 2],pr_plots[::100, 0]);
    
    ax2.scatter(pr_plots[::100, 4], pr_plots[::100, 5],pr_plots[::100, 3]);
    
    ax3.scatter(pr_plots[::100, 7], pr_plots[::100, 8],pr_plots[::100, 6]);
    
    #ax3.scatter(pr_plots[:, 3]+pr_plots[:500, 0], pr_plots[:500, 4]+pr_plots[:500, 1],pr_plots[:500, 5]+pr_plots[:500, 2]);
    
    plt.figure();
    
    plt.plot(spd_est);

    #spd_est=[];

nanVec = np.sum(np.isnan(latlong_mat), axis=1)==0;

print("Nan values:", np.sum(nanVec==False));


print(np.mean(pr_plots, axis=0), "mean projection errs");

sub1['LatitudeDegrees'][nanVec] = latlong_mat[nanVec, 0];

sub1['LongitudeDegrees'][nanVec] = latlong_mat[nanVec, 1];

sub1["UnixTimeMillis"] = sub1["UnixTimeMillis"].to_numpy()
        
    #plt.figure();
    
    #plt.plot(train_err[::50]);
    
print(np.sum(sub1['LatitudeDegrees'] == sub2['LatitudeDegrees']), np.shape(sub1));
    
#pd.DataFrame(sub1).to_csv('sub1.csv', index=False);

sub1.to_csv('/kaggle/working/sub1.csv', index = False);

##

sc = list(sub1.columns);

#sc[0] = "phone";

sub2_v = np.ones(sub1.shape[0], bool);

sub2_v[46214] = False;
sub2_v[50267] = False;
sub2_v[50268] = False;

print("nans:", latlong_mat[sub2_v==False, :])

latlong_mat[46214,:] = latlong_mat[46213,:]
latlong_mat[50267,:] = latlong_mat[50266,:]
latlong_mat[50268,:] = latlong_mat[50267,:]

sub2['LatitudeDegrees'] = latlong_mat[:, 0];
sub2['LongitudeDegrees'] = latlong_mat[:, 1];

#sub2['LatitudeDegrees'][sub2_v] = latlong_mat[sub2_v, 0];
#sub2['LongitudeDegrees'][sub2_v] = latlong_mat[sub2_v, 1];


#sub2['LatitudeDegrees'][:46210] = latlong_mat[:46210, 0];
#sub2['LongitudeDegrees'][:46210] = latlong_mat[:46210, 1];

#sub2['LatitudeDegrees'][50275:] = latlong_mat[50275:, 0];
#sub2['LongitudeDegrees'][50275:] = latlong_mat[50275:, 1];

#sub2['LatitudeDegrees'][46220:50265] = latlong_mat[46220:50265, 0];
#sub2['LongitudeDegrees'][46220:50265] = latlong_mat[46220:50265, 1];


pd.DataFrame(np.array(sub2), columns=sc, index=None).to_csv("/kaggle/working/sub2.csv", index=False);
    #sub1.save_csv()

    
#f = open("/kaggle/working/sub1.csv", "r");


#fig.savefig("sat_positions.jpg");

#fig2.savefig("sat_vel.jpg");

#fig3.savefig("ground_positions.jpg");

#ax = fig.add_subplot(projection='3d')
#ax.scatter(t2["WlsPositionYEcefMeters"][::250],t2["WlsPositionXEcefMeters"][::250],t2["WlsPositionZEcefMeters"][::250]);

##pd.DataFrame(,)
# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using "Save & Run All" 
# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session