{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":101849,"databundleVersionId":13093295,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13139844,"sourceType":"datasetVersion","datasetId":8324725},{"sourceId":13140115,"sourceType":"datasetVersion","datasetId":8324909},{"sourceId":13144472,"sourceType":"datasetVersion","datasetId":8324898},{"sourceId":13229891,"sourceType":"datasetVersion","datasetId":8385999},{"sourceId":13229910,"sourceType":"datasetVersion","datasetId":8386010},{"sourceId":13229933,"sourceType":"datasetVersion","datasetId":8386026}],"dockerImageVersionId":31089,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Version 7 of this notebook is the one that was submitted to the competition.  All subsequent versions aim to improve my understanding of an apparent platform dependence that cropped up right before the submission deadline and the level of its importance to the difference between the scores I saw in local CV and the ones I see in the public/private test sets.","metadata":{}},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport os\n\n#doesn't actually seem to help with non-determinism\n#os.environ[\"MKL_NUM_THREADS\"] = \"1\"\n#os.environ[\"OMP_NUM_THREADS\"] = \"1\"   # often useful too\n#os.environ[\"NUMEXPR_NUM_THREADS\"] = \"1\"\n#os.environ[\"OPENBLAS_NUM_THREADS\"] = \"1\"\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport scipy\nimport itertools\nfrom scipy.optimize import curve_fit\nimport torch\nimport re\nimport time as time_lib\nimport collections\nfrom collections import OrderedDict\nimport multiprocessing as mp\n#!pip install astropy --target=/kaggle/working/\n#!pip install PyWavelets==1.7.0 --target=/kaggle/working/\n!ls -lrtha /kaggle/working/\n\nfrom astropy.stats import sigma_clip\nimport pywt\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n#for dirname, _, filenames in os.walk('/kaggle/input'):\n#    for filename in filenames:\n#        print(os.path.join(dirname, filename))\nnp.random.seed(334258)\ntorch.manual_seed(0)\ntorch.use_deterministic_algorithms(True)\ntorch.backends.cudnn.benchmark = False\ntorch.utils.deterministic.fill_uninitialized_memory=True\n\nprint(\"numpy version: \"+str(np.version.version),flush=True)\nprint(\"scipy version: \"+str(scipy.__version__),flush=True)\nprint(\"torch version: \"+str(torch.__version__),flush=True)\n#!pip freeze\n\nif torch.cuda.is_available():\n  device=torch.device(\"cuda\")\nelse:\n  device=torch.device(\"cpu\")\nDEVICE=device\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n#for dirname, _, filenames in os.walk('/kaggle/input'):\n#    for filename in filenames:\n#        print(os.path.join(dirname, filename))\n\n# 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\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session\n\n#print(\"numpy show config:\",flush=True)\n#np.show_config()\n#from iminuit import Minuit\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-09T03:06:11.797773Z","iopub.execute_input":"2025-10-09T03:06:11.798271Z","iopub.status.idle":"2025-10-09T03:06:16.935708Z","shell.execute_reply.started":"2025-10-09T03:06:11.798246Z","shell.execute_reply":"2025-10-09T03:06:16.9348Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# control switches and global constants","metadata":{}},{"cell_type":"code","source":"#control switches and gloabal constants\n#dataset=\"train\"  \ndataset=\"test\"\nmax_num_train_samples=3  #in the kernel environment, runs on the train set are just for debugging.  So only run this many.\ntrain_sample_start=167 #884\n\nn_neldermead=0\nn_neldermead_max=250\n\n\n#even though it's called \"cpu_fit\", some implementations have made use of the gpu.  Choose this switch based on speed.\ncpu_fit_device=torch.device('cpu')\n#cpu_fit_device=device\n#hess_device=torch.device('cpu') #perhaps the hess computations for airs are also so unable to leverage gpu parallelism that it makes sense to run them on the cpu too?\nhess_device=device   #...no, I don't see a significant difference.\n\ndo_linear_corr=True\nch0_scistart=39\nch0_sciend=321\nch0_nnstart=ch0_scistart-25   #36\nch0_nnend=ch0_sciend+25   #324\n#nrebin=None  #25  #75\nfit_bin_width=15  #25  #75\n\nNPARAM=4\nNPARAM_MAX=4\ndrop_shoulders=25  #25\nnormalize_batch_fits=True\nMAX_LDC_COEF=0.5\n\n#before fits, errors are smoothed by averaging over neighboring bins. \ndo_error_smoothing=True  \nerror_smoothing_scalefac=1  \nSMOOTH_SIZE=2\n\nn_hidden=64\nbatch_size=128\nREBINNINGS=[3,7,47]  \n\nwavelet = \"cgau1\"\nwidths = np.geomspace(1, 1024, num=10)\ntime = np.linspace(0, 1, 5625)\nsampling_period = np.diff(time).mean()\n\nturn_on_width=2\n\n#when combining measurements from two visits, should we take the difference between the two means as a systematic uncertainty?\ndifference_is_uncert=True\n\n#sometimes fits fail to converge well enough to compute reliable uncertainties using the Hessian. \n#This is hard to avoid given the runtime constraints in this contest.  It's a challenge for the FGS \n#fit in particular, since this analysis uses the FGS fit to pin down certain fit parameters (like Tcenter)\n#which are then held constant in the AIRS fits that come after.  Thankfully, it is usually the case that the\n#parameter values are good enough to get the job done, even if the Hessian computation fails.  So in these \n#cases, default to the average fit uncertainties from the training set\n#will load these from a file generated during skim for the final run, but for now, temporary values extracted by hand from an older skim:\n#avg_hess_uncerts=[0.00010687391152194368,0.06983016436985619,0.11065483080895322,0.05325828968780703,4.6466757371942e-05,0.06129590380754077,0.027488655060930314,24.79859075210302,0.0016164013792447425]\navg_hess_uncerts=pd.read_csv(\"/kaggle/input/ariel2025-average-fit-errors-v19-debugncg/average_fit_errors_v19_debugNCG.csv\")\n\n#simple mean and sigma of the training labels, used as default \n#values in case a prediction comes up with nan or inf\nnaive_mean=0.014689019532534075\nnaive_sigma=0.01066133533197834 \n\n#placeholders:  these get populated at the start of main\nnn_input_var_min=None\nnn_input_var_max=None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-09T03:06:16.937363Z","iopub.execute_input":"2025-10-09T03:06:16.937788Z","iopub.status.idle":"2025-10-09T03:06:16.96515Z","shell.execute_reply.started":"2025-10-09T03:06:16.937765Z","shell.execute_reply":"2025-10-09T03:06:16.964448Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# functions used in the (preliminary) detection of transit start/end","metadata":{}},{"cell_type":"code","source":"def gauss_func(x,mu,sig,norm):\n  return (norm/(sig*(2*np.pi)**0.5))*np.exp(-0.5*(x-mu)**2/sig**2)\n\ndef find_peak(x,y,size=10,fit_window=1000):\n  assert len(x.shape)==1\n  assert len(y.shape)==1\n  assert x.shape==y.shape\n  if not isinstance(x,np.ndarray):\n    x=x.numpy()\n  if not isinstance(y,np.ndarray):\n    y=y.numpy()\n  smoothed=scipy.ndimage.median_filter(y,size=size)\n  initial_guess=np.argmax(smoothed)\n  start=max(0,initial_guess-fit_window//2)\n  end=min(x.shape[0]-1,initial_guess+fit_window//2)\n  xdata=x[start:end]\n  ydata=y[start:end]\n  lower=[x[start],0.1,0]\n  upper=[x[end],x[end]-x[start],np.inf]\n  popt,pcov=curve_fit(gauss_func,x,y,bounds=(lower,upper))\n  return popt,start,end\n\ndef weighted_average(coords):\n  wnorm=sum([1/c[1] for c in coords])\n  weights=[(1/c[1])/wnorm for c in coords]\n  mean=sum([w*c[0] for w,c in zip(weights,coords)])\n  sig=sum([(w*c[1])**2 for w,c in zip(weights,coords)])**0.5\n  return mean,sig\n\n\ndef peak_average(coords):\n  #just a simple weighted average, but remove outliers\n  assert len(coords)>0\n  if len(coords)==1:\n    return tuple(coords[0])\n  mean=sum([c[0] for c in coords])/len(coords)\n  width=(sum([c[1]**2 for c in coords])/len(coords))**0.5\n  residuals=[(c[0]-mean)/width for c in coords]\n  keep=[abs(r)<2 for r in residuals]\n  if all(keep):\n    return weighted_average(coords)\n  else:\n    return peak_average([c for c,k in zip(coords,keep) if k])\n\ndef find_transit_region(ch0_signal,is_subtracted=False):\n\n    if is_subtracted:\n      white=ch0_signal\n    else:\n      white=torch.sum(ch0_signal,dim=1)\n\n    norm=(torch.mean(white[:,0:50],dim=-1,keepdim=True)+torch.mean(white[:,-50:],dim=-1,keepdim=True))/2\n    white=white/norm\n    white=white-1\n\n    cwtmatr, freqs = pywt.cwt(torch.squeeze(white).cpu().numpy(), widths, wavelet, sampling_period=sampling_period)\n    cwtmatr = np.abs(cwtmatr[:-1, :-1])\n\n    firstpeak_coords=[]\n    secondpeak_coords=[]\n    for row in [1,2,3,4,5]:\n      scalslice=np.expand_dims(cwtmatr[-row,:],axis=0)\n      x=np.linspace(0,scalslice.shape[-1],num=scalslice.shape[-1])\n      popt_firstpeak,start,end=find_peak(x,scalslice[0,:],size=10,fit_window=1500)\n\n      scalslice_masked=np.array(scalslice)\n      scalslice_masked[0,start:end]=0\n      popt_secondpeak,start2,end2=find_peak(x,scalslice_masked[0,:],size=10,fit_window=1500)\n\n      #reminder: popt_firstpeak is a list like [1820.31765939  183.7399781     5.72666547], i.e., fitted mu, sigma, and norm of gaussian\n      if popt_firstpeak[0]>popt_secondpeak[0]:\n        popt_firstpeak,popt_secondpeak=popt_secondpeak,popt_firstpeak\n      firstpeak_coords.append(popt_firstpeak[0:2])\n      secondpeak_coords.append(popt_secondpeak[0:2])\n\n    firstpeak_mean,firstpeak_sig=peak_average(firstpeak_coords)\n    secondpeak_mean,secondpeak_sig=peak_average(secondpeak_coords)\n\n    print(\"first-pass transit region estimates:  \"+str(firstpeak_mean)+\"+/-\"+str(firstpeak_sig)+\" and \"+str(secondpeak_mean)+\"+/-\"+str(secondpeak_sig),flush=True)\n\n    if is_subtracted:\n      return firstpeak_mean,firstpeak_sig,secondpeak_mean,secondpeak_sig,cwtmatr\n    else:\n      gap1=(int((firstpeak_mean-firstpeak_sig*turn_on_width)/fit_bin_width)*fit_bin_width,\n            fit_bin_width*(1+int((firstpeak_mean+firstpeak_sig*turn_on_width)/fit_bin_width)))\n      gap2=(int((secondpeak_mean-secondpeak_sig*turn_on_width)/fit_bin_width)*fit_bin_width,\n            fit_bin_width*(1+int((secondpeak_mean+secondpeak_sig*turn_on_width)/fit_bin_width)))\n\n      white_x=np.arange(0,1,1./white.shape[-1])\n      white_x_nogap=[white_x[0:gap1[0]],\n                     white_x[gap1[1]:gap2[0]],\n                     white_x[gap2[1]:]]\n      white_x_nogap=np.concatenate(white_x_nogap,axis=-1)\n\n      white_nogap=[white[:,0:gap1[0]],\n                   white[:,gap1[1]:gap2[0]],\n                   white[:,gap2[1]:]]\n      white_nogap=torch.cat(white_nogap,dim=-1).cpu().numpy()\n\n      white_bins=np.reshape(white_nogap,tuple(list(white_nogap.shape)[0:-1]+[-1,fit_bin_width]))\n      white_y=np.mean(white_bins,axis=-1)\n      white_yerr=np.std(white_bins,axis=-1)/fit_bin_width**0.5\n\n      if do_error_smoothing:\n        #smooth out the errors so that fluctuations don't give us very different errors in neighboring bins\n        pad_size=tuple([(0,0) for idim in range(len(white_yerr.shape)-1)]+[(2,2)])\n        white_yerr=np.lib.stride_tricks.sliding_window_view(np.pad(white_yerr,pad_size,'edge'),5,axis=-1)\n        white_yerr=np.mean(white_yerr,axis=-1)*error_smoothing_scalefac\n\n      white_x_binned=np.mean(np.reshape(white_x_nogap,(-1,fit_bin_width)),axis=-1)\n\n\n      in_transit=np.concatenate([np.zeros_like(white_x[0:gap1[0]]),\n                np.ones_like(white_x[gap1[1]:gap2[0]]),\n                np.zeros_like(white_x[gap2[1]:])],axis=0)\n      in_transit=np.reshape(in_transit,(-1,fit_bin_width))\n      in_transit=np.min(in_transit,axis=-1)\n\n      init_norm=np.mean(white_bins)\n      init_data=np.array(tuple([0 for i in range(NPARAM)]+[init_norm]),dtype=np.float64)\n\n      data=(white_x_binned,white_y[0,:],white_yerr[0,:],in_transit)\n      res=scipy.optimize.minimize(light_curve_chisq_simplified,init_data,  #first 0 is for offset, not a polynomial param\n                                bounds=None,\n                                args=data,jac=True)\n\n      coeffs=res.x[1:]\n      coeffs=coeffs[::-1]\n      fit_y=np.zeros_like(white_x)\n      for i in range(len(coeffs)):\n        fit_y+=coeffs[i]*white_x**i\n\n      white_sub=white.cpu().numpy()-fit_y\n      try:\n        firstpeak_mean2,firstpeak_sig2,secondpeak_mean2,secondpeak_sig2,cwtmatr2=find_transit_region(torch.tensor(white_sub),is_subtracted=True)\n        avg=(firstpeak_sig+secondpeak_sig)/2\n        avg2=(firstpeak_sig2+secondpeak_sig2)/2\n        if avg2<=avg:\n          return firstpeak_mean2,firstpeak_sig2,secondpeak_mean2,secondpeak_sig2,cwtmatr2\n        else:\n          #print(\"second-pass transit region estimates came out bigger than first-pass; rolling back to first-pass\",flush=True)\n          return firstpeak_mean,firstpeak_sig,secondpeak_mean,secondpeak_sig,cwtmatr\n\n      except Exception as e:\n        #print(\"failed to find second-pass transit region estimates: \"+str(e),flush=True)\n        return firstpeak_mean,firstpeak_sig,secondpeak_mean,secondpeak_sig,cwtmatr\n\ndef light_curve_chisq_simplified(params,*args):\n                                \n  #reminder: ch0_signal.shape=torch.Size([1, 288, 5625])\n  #but we are fitting one wavelength at a time so that the optimizer \n  #doesn't get screwed up exploring stupid cross-wavelength correlations.\n  #so we expect arguments to have the following shapes:\n  #  - obs_x, obs_y, obs_yerr, and in_transit:  (5625,)\n  #  - offset and polynomial coefficients: scalar\n        \n  assert len(args)==4\n  obs_x=torch.tensor(args[0])\n  obs_y=torch.tensor(args[1])\n  obs_yerr=torch.tensor(args[2])\n  in_transit=torch.tensor(args[3])\n          \n  assert len(obs_x.shape)==1\n  assert obs_x.shape==obs_y.shape\n  assert obs_x.shape==obs_yerr.shape\n  assert obs_x.shape==in_transit.shape\n  params=torch.tensor(params)\n  assert len(params.shape)==1\n  params.requires_grad=True\n\n  offset=params[0]\n  coeffs=params[1:]\n  coeffs=torch.flip(coeffs,dims=(0,))\n  fit_y=torch.zeros_like(obs_x)\n  for i in range(len(coeffs)):\n    fit_y+=coeffs[i]*obs_x**i\n\n  fit_y=fit_y-torch.where(in_transit>=1,offset,0)\n\n  chisq=(torch.abs(obs_y-fit_y))/obs_yerr  \n  chisq=chisq**2\n  chisq=torch.sum(chisq)\n\n  chisq.backward()\n  grad=params.grad.detach().numpy()\n  return chisq.detach().numpy(),grad\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-09T03:06:16.965931Z","iopub.execute_input":"2025-10-09T03:06:16.966127Z","iopub.status.idle":"2025-10-09T03:06:16.990235Z","shell.execute_reply.started":"2025-10-09T03:06:16.966111Z","shell.execute_reply":"2025-10-09T03:06:16.989662Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# functions for loading and calibrating the contest data","metadata":{}},{"cell_type":"code","source":"def load_metadata():\n  if dataset==\"train\":\n    fname=\"/kaggle/input/ariel-data-challenge-2025/train.csv\"\n  else:\n    fname=\"/kaggle/input/ariel-data-challenge-2025/sample_submission.csv\"\n\n  ss = pd.read_csv(fname) #'/kaggle/input/ariel-data-challenge-2025/sample_submission.csv')\n  planets=[str(int(p)) for p in ss[\"planet_id\"].to_list()]\n\n  if dataset==\"train\":\n    planets=planets[train_sample_start:train_sample_start+max_num_train_samples]\n      \n  #all planets use the same adc info in the 2025 version of this contest\n  #FGS1_adc_offset,FGS1_adc_gain,AIRS-CH0_adc_offset,AIRS-CH0_adc_gain\n  #-1000.0,0.4369,-1000.0,0.4369\n  df_adc=pd.read_csv(\"/kaggle/input/ariel-data-challenge-2025/adc_info.csv\")\n  df_axis=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/axis_info.parquet\")\n  integration_time=df_axis[\"AIRS-CH0-integration_time\"].dropna().to_numpy()\n  cumulative_time=np.cumsum(integration_time,axis=-1)\n  cumulative_time_doublesamp=cumulative_time[1::2]\n\n  data_dict=dict()\n  data_dict[\"integration_time\"]=integration_time\n  data_dict[\"cumulative_time\"]=cumulative_time\n  data_dict[\"cumulative_time_doublesamp\"]=cumulative_time_doublesamp\n  data_dict[\"index\"]=planets\n    \n  for iplanet,planet in enumerate(planets):\n    row=df_adc.iloc[0]  #df_adc[\"planet_id\"]==int(planet)]\n    fgs_offset=row[\"FGS1_adc_offset\"]\n    fgs_gain=row[\"FGS1_adc_gain\"]\n    ch0_offset=row[\"AIRS-CH0_adc_offset\"]\n    ch0_gain=row[\"AIRS-CH0_adc_gain\"]\n    star=int(planet)\n    data_dict[planet]=(ch0_offset,ch0_gain,fgs_offset,fgs_gain,star)\n\n  return data_dict\n\ndef mask_hot_dead(signal, dead, dark):\n    hot = sigma_clip(\n        dark, sigma=5, maxiters=5\n    ).mask\n    hot = np.tile(hot, (signal.shape[0], 1, 1))\n    dead = np.tile(dead, (signal.shape[0], 1, 1))\n\n    hot=torch.tensor(hot.astype(np.int32)).to(DEVICE)\n    dead=torch.tensor(dead.astype(np.int32)).to(DEVICE)\n    signalmask=torch.maximum(hot,dead)\n    return signalmask\n\ndef clean_dark(signal, signalmask, dead, dark, dt):\n    dark = np.ma.masked_where(dead, dark)\n    dark = np.tile(dark, (signal.shape[0], 1, 1))\n    darkmask = np.tile(dead, (signal.shape[0],1,1))\n    dark=torch.tensor(dark).to(DEVICE)\n    darkmask=torch.tensor(darkmask).to(DEVICE)\n    mask=torch.maximum(darkmask.float(),signalmask.float())\n    if not torch.is_tensor(dt):\n      dt=torch.tensor(dt)\n\n    signal=torch.where(mask>=1,signal,signal-dark*torch.unsqueeze(torch.unsqueeze(dt.to(DEVICE),dim=-1),dim=-1))\n    return signal,mask\n\ndef correct_flat_field(flat,dead, signal,signalmask):\n    flat = np.ma.masked_where(dead, flat)\n    flat = np.tile(flat, (signal.shape[0], 1, 1))\n    flatmask=torch.tensor(flat.mask).to(DEVICE)\n    flat=torch.tensor(flat.data).to(DEVICE)\n    mask=torch.maximum(flatmask.float(),signalmask.float())\n    signal=torch.where(mask>=1,signal,signal/flat)\n    return signal,mask\n\ndef linear_correction(signal,lincorr):\n  signal_lincorr=torch.zeros_like(signal)\n  for ic in range(lincorr.shape[0]):\n    signal_lincorr=signal_lincorr+torch.einsum('jk,ijk->ijk',lincorr[ic,:,:],signal**ic)\n\n  return signal_lincorr\n\n\ndef load_planet(planet,ch0_offset,ch0_gain,fgs_offset,fgs_gain,integration_time,visit):\n  #print(\"hello from load_planet for planet \"+str(planet),flush=True)\n  #print(\"DEVICE=\"+str(DEVICE),flush=True)\n  assert DEVICE is not None\n\n  #print(\"reading fgs dataframe from filename /kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/FGS1_signal_\"+str(visit)+\".parquet\",flush=True)\n  df_fgs=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/FGS1_signal_\"+str(visit)+\".parquet\")\n  #print(\"got it\",flush=True)\n  fgs_signal=torch.tensor(df_fgs.to_numpy().astype(np.float32)).to(DEVICE)\n  fgs_signal = fgs_signal.reshape((fgs_signal.shape[0], 32, 32))\n  fgs_signal_corr=fgs_signal/fgs_gain+fgs_offset\n  dt_fgs1 = torch.ones(len(fgs_signal_corr)).to(DEVICE)*0.1\n  dt_fgs1[1::2] += 0.1\n\n  #dead/dark/flat correction\n  fgs_flat=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/FGS1_calibration_\"+str(visit)+\"/flat.parquet\").values.astype(np.float64).reshape((32, 32))\n  fgs_dark=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/FGS1_calibration_\"+str(visit)+\"/dark.parquet\").values.astype(np.float64).reshape((32, 32))\n  fgs_dead=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/FGS1_calibration_\"+str(visit)+\"/dead.parquet\").values.astype(np.float64).reshape((32, 32))\n    \n  fgs_corr_mask = mask_hot_dead(fgs_signal_corr, fgs_dead, fgs_dark)\n    \n  df_fgs_lin=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/FGS1_calibration_\"+str(visit)+\"/linear_corr.parquet\")\n  #print(\"df_fgs_lin read ok\",flush=True)\n    \n  fgs_lin=torch.tensor(df_fgs_lin.to_numpy().astype(np.float32)).to(DEVICE)\n  fgs_coeffs=fgs_lin.reshape((6,32,32))\n  if do_linear_corr:\n    fgs_signal_corr=linear_correction(fgs_signal_corr,fgs_coeffs)\n\n  fgs_corr,fgs_corr_mask = clean_dark(fgs_signal_corr,fgs_corr_mask, fgs_dead, fgs_dark,dt_fgs1)\n  fgs_corr = fgs_corr[1::2,:,:] - fgs_corr[::2,:,:]\n  fgs_corr_mask = torch.maximum(fgs_corr_mask[1::2,:,:],fgs_corr_mask[::2,:,:])\n  fgs_corr,fgs_corr_mask = correct_flat_field(fgs_flat,fgs_dead, fgs_corr,fgs_corr_mask)\n\n  #print(\"fgs_corr.shape=\"+str(fgs_corr.shape),flush=True)\n  fgs_corr_proj1=torch.clip(torch.sum(torch.where(fgs_corr_mask<=0,fgs_corr,0),dim=-1),0,1e30)\n  fgs_norm_proj1=torch.sum(fgs_corr_proj1,dim=-1,keepdim=True)\n  fgs_probs_proj1=fgs_corr_proj1/fgs_norm_proj1\n  fgs_corr_proj2=torch.clip(torch.sum(torch.where(fgs_corr_mask<=0,fgs_corr,0),dim=-2),0,1e30)\n  fgs_norm_proj2=torch.sum(fgs_corr_proj2,dim=-1,keepdim=True)\n  fgs_probs_proj2=fgs_corr_proj2/fgs_norm_proj2\n  x=torch.unsqueeze(torch.arange(0,32).to(DEVICE),dim=0)\n  fgs_mean_proj1=torch.sum(x*fgs_probs_proj1,dim=-1,keepdim=True)\n  #print(\"x.shape=\"+str(x.shape)+\", fgs_mean_proj1.shape=\"+str(fgs_mean_proj1.shape),flush=True)\n\n  fgs_std_proj1=torch.sum(fgs_probs_proj1*(x-fgs_mean_proj1)**2,dim=-1)\n  fgs_mean_proj2=torch.sum(x*fgs_probs_proj2,dim=-1,keepdim=True)\n  fgs_std_proj2=torch.sum(fgs_probs_proj2*(x-fgs_mean_proj2)**2,dim=-1)\n\n  fgs_mean_proj1=torch.mean(torch.reshape(fgs_mean_proj1,(-1,12*fit_bin_width)),dim=-1)\n  fgs_std_proj1=torch.mean(torch.reshape(fgs_std_proj1,(-1,12*fit_bin_width)),dim=-1)\n  fgs_mean_proj2=torch.mean(torch.reshape(fgs_mean_proj2,(-1,12*fit_bin_width)),dim=-1)\n  fgs_std_proj2=torch.mean(torch.reshape(fgs_std_proj2,(-1,12*fit_bin_width)),dim=-1)\n\n  #simple outlier exclusion -- make a sliding window in time and when you find a point >3sigma away from \n  #the mean in the sliding window, replace it with the mean.  Should get rid of stuff like cosmic rays.\n  #The \"slide\" in these variable names refers to \"sliding window\"\n  fgs_slide=fgs_corr.unfold(dimension=0,size=50,step=1)\n  fgs_slide_mean=torch.mean(fgs_slide,dim=-1)\n  fgs_slide_std=torch.std(fgs_slide,dim=-1)\n  fgs_slide_mean=torch.transpose(fgs_slide_mean,0,2)\n  fgs_slide_std=torch.transpose(fgs_slide_std,0,2)\n  fgs_slide_mean=torch.nn.functional.pad(fgs_slide_mean,(25,24),mode='replicate')\n  fgs_slide_std=torch.nn.functional.pad(fgs_slide_std,(25,24),mode='replicate')\n  fgs_slide_mean=torch.transpose(fgs_slide_mean,0,2)\n  fgs_slide_std=torch.transpose(fgs_slide_std,0,2)\n\n  dev=(fgs_corr-fgs_slide_mean)/fgs_slide_std\n  fgs_corr=torch.where(torch.abs(dev)<3,fgs_corr,fgs_slide_mean)\n\n  fgs_corr=torch.unsqueeze(torch.sum(torch.where(fgs_corr_mask<=0,fgs_corr,0),axis=(-1,-2)),dim=0)\n  fgs_nmask=torch.sum(fgs_corr_mask)\n\n  #print(\"reading ch0\",flush=True)\n  df_ch0=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/AIRS-CH0_signal_\"+str(visit)+\".parquet\")\n  ch0_signal=torch.tensor(df_ch0.to_numpy().astype(np.float32),dtype=torch.float32).to(DEVICE)\n  ch0_signal=ch0_signal.reshape(-1,32,356)\n  ch0_signal=ch0_signal[:,:,ch0_nnstart:ch0_nnend]\n  #ADC correction\n  ch0_signal_corr=ch0_signal/ch0_gain+ch0_offset\n\n  #do we need to scale these by collection time?  https://www.kaggle.com/competitions/ariel-data-challenge-2024/discussion/528066\n  #dead/dark/flat correction\n  df_ch0_flat=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/AIRS-CH0_calibration_\"+str(visit)+\"/flat.parquet\")\n  ch0_flat=df_ch0_flat.to_numpy()[:,ch0_nnstart:ch0_nnend]\n\n  df_ch0_dark=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/AIRS-CH0_calibration_\"+str(visit)+\"/dark.parquet\")\n  ch0_dark=df_ch0_dark.to_numpy()[:,ch0_nnstart:ch0_nnend]\n\n  df_ch0_dead=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/AIRS-CH0_calibration_\"+str(visit)+\"/dead.parquet\")\n  ch0_dead=df_ch0_dead.to_numpy()[:,ch0_nnstart:ch0_nnend]\n\n  #print(\"calling ch0 mask_hot_dead\",flush=True)\n  #mask hot/dead\n  ch0_signal_corr_mask=mask_hot_dead(ch0_signal_corr,ch0_dead,ch0_dark)\n\n  df_ch0_lin=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/AIRS-CH0_calibration_\"+str(visit)+\"/linear_corr.parquet\")\n  ch0_lin=torch.tensor(df_ch0_lin.to_numpy().astype(np.float32)).to(DEVICE)\n  ch0_coeffs=ch0_lin.reshape((6,32,356))\n  ch0_coeffs=ch0_coeffs[:,:,ch0_nnstart:ch0_nnend]\n\n  if do_linear_corr:\n    ch0_signal_lincorr=linear_correction(ch0_signal_corr,ch0_coeffs)\n    ch0_lincorr_shift=ch0_signal_lincorr-ch0_signal_corr\n    ch0_signal_corr=ch0_signal_lincorr\n\n  #dark current subtraction\n  dt_airs=np.array(integration_time)\n  dt_airs[1::2]+=0.1\n  ch0_signal_corr,sch0_signal_corr_mask=clean_dark(ch0_signal_corr,ch0_signal_corr_mask,ch0_dead,ch0_dark,dt_airs)  #integration_time)\n\n  #double-sampling correction\n  ch0_dsamp=ch0_signal_corr[::2,:,:]\n  ch0_signal_corr=ch0_signal_corr[1::2,:,:]-ch0_signal_corr[::2,:,:]\n  ch0_signal_corr_mask=torch.maximum(ch0_signal_corr_mask[1::2,:,:],ch0_signal_corr_mask[::2,:,:])\n\n  #print(\"calling ch0 flat field correction\",flush=True)\n  #flat field correction\n  ch0_signal_corr,ch0_signal_corr_mask=correct_flat_field(ch0_flat,ch0_dead,ch0_signal_corr,ch0_signal_corr_mask)\n\n  #read noise\n  #df_ch0_read=pd.read_parquet(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/AIRS-CH0_calibration_\"+str(visit)+\"/read.parquet\")\n  #ch0_read=df_ch0_read.to_numpy()[:,ch0_nnstart:ch0_nnend] #before that last trim, ch0_read.shape=(32, 356), dtype=float64\n\n  ch0_norm=torch.sum(torch.clip(ch0_signal_corr,0,1e30),dim=-2,keepdim=True)\n  ch0_probs=ch0_signal_corr/ch0_norm\n  x=torch.unsqueeze(torch.unsqueeze(torch.arange(0,32).to(DEVICE),dim=0),dim=-1)\n  ch0_mean=torch.sum(x*ch0_probs,dim=-2,keepdim=True)\n  ch0_std=torch.sum(ch0_probs*(x-ch0_mean)**2,dim=-2)\n\n  ch0_mean=torch.mean(torch.reshape(torch.transpose(ch0_mean,0,1),(ch0_mean.shape[-1],-1,fit_bin_width)),dim=-1)\n  ch0_std=torch.mean(torch.reshape(torch.transpose(ch0_std,0,1),(ch0_std.shape[-1],-1,fit_bin_width)),dim=-1)\n\n  ch0_signal_corr=torch.sum(torch.where(ch0_signal_corr_mask<=0,ch0_signal_corr,0),axis=1)  #ignore spatial dim for now -- shape=(n_timesteps,282)\n  ch0_signal_corr=torch.transpose(ch0_signal_corr,0,1)  #shape=(282,n_timesteps), so that the time series is an event dim\n  ch0_nmask=torch.sum(ch0_signal_corr_mask,dim=(0,1))\n\n  #print(\"before do_linear_corr if block, so far so good\",flush=True)\n  if do_linear_corr:\n    #print(\"ch0_lincorr_shift.shape=\"+str(ch0_lincorr_shift.shape),flush=True)\n    ch0_lincorr_shift=torch.sum(ch0_lincorr_shift,axis=1)\n    #print(\"ch0_dsamp.shape=\"+str(ch0_dsamp.shape),flush=True)\n    ch0_dsamp=torch.sum(ch0_dsamp,axis=1)\n\n    return torch.unsqueeze(ch0_signal_corr,dim=0),torch.unsqueeze(fgs_corr,dim=0), ch0_lincorr_shift, ch0_dsamp, fgs_nmask,ch0_nmask,fgs_mean_proj1, fgs_std_proj1, fgs_mean_proj2, fgs_std_proj2, ch0_mean, ch0_std\n  else:\n    return torch.unsqueeze(ch0_signal_corr,dim=0),torch.unsqueeze(fgs_corr,dim=0), None, ch0_dsamp, fgs_nmask,ch0_nmask,fgs_mean_proj1, fgs_std_proj1, fgs_mean_proj2, fgs_std_proj2, ch0_mean, ch0_std\n\ndef load_star_params():\n  \n  \"\"\"\n  head /ariel_data/train_star_info.csv\n  planet_id,Rs,Ms,Ts,Mp,e,P,sma,i\n  34983,1.155435480707952,1.062960837903184,5577.006645157513,0.6949463499730458,0.0,3.305588751623328,8.550785903749272,89.1507586203837\n  1873185,1.813230199682509,1.370450683964498,6216.229756270119,0.6108447312470062,0.0,6.352659805895124,9.55338410018746,88.70151407048026\n  3849793,0.6538067391809703,0.667352384889845,4968.477185692076,1.5291995382480204,0.0,5.522797615237956,15.285660679929052,89.1341766239202\n  \"\"\"\n\n  df_starinfo=pd.read_csv(\"/kaggle/input/ariel-data-challenge-2025/train_star_info.csv\")\n  var_min=df_starinfo.min()\n  var_max=df_starinfo.max()\n\n  if dataset==\"test\":\n    df_starinfo=pd.read_csv(\"/kaggle/input/ariel-data-challenge-2025/test_star_info.csv\")\n\n  #print(\"var_min=\"+str(var_min),flush=True)\n  #print(\"var_max=\"+str(var_max),flush=True)\n    \n  star_params=dict()\n  star_params_norm=dict()\n  for istar in range(len(df_starinfo)):\n    planet_id=str(int(df_starinfo.iloc[istar][\"planet_id\"]))\n    Rs=df_starinfo.iloc[istar][\"Rs\"]\n    Ms=df_starinfo.iloc[istar][\"Ms\"]\n    Ts=df_starinfo.iloc[istar][\"Ts\"]\n    Mp=df_starinfo.iloc[istar][\"Mp\"]\n    e=df_starinfo.iloc[istar][\"e\"]\n    P=df_starinfo.iloc[istar][\"P\"]*24  #days-->hours      #*24*60*60  #days-->seconds\n    sma=df_starinfo.iloc[istar][\"sma\"]\n    i=df_starinfo.iloc[istar][\"i\"]*np.pi/180\n    star_params[planet_id]=(Rs,Ms,Ts,Mp,e,P,sma,i)\n    b=max(0,min(1,sma*np.cos(i)))  #...and e is always 0 in this dataset\n\n    Rsnorm=(Rs-var_min[\"Rs\"])/(var_max[\"Rs\"]-var_min[\"Rs\"])\n    Msnorm=(Ms-var_min[\"Ms\"])/(var_max[\"Ms\"]-var_min[\"Ms\"])\n    Tsnorm=(Ts-var_min[\"Ts\"])/(var_max[\"Ts\"]-var_min[\"Ts\"])\n    Mpnorm=(Mp-var_min[\"Mp\"])/(var_max[\"Mp\"]-var_min[\"Mp\"])\n    #enorm=(e-var_min[\"e\"])/(var_max[\"e\"]-var_min[\"e\"])  #<---always zero, I think, so norm comes out nan\n    Pnorm=(P-var_min[\"P\"]*24)/(var_max[\"P\"]*24-var_min[\"P\"]*24)\n    smanorm=(sma-var_min[\"sma\"])/(var_max[\"sma\"]-var_min[\"sma\"])\n    inorm=(i-var_min[\"i\"]*np.pi/180)/(var_max[\"i\"]*np.pi/180-var_min[\"i\"]*np.pi/180)\n    star_params_norm[planet_id]=(Rsnorm,Msnorm,Tsnorm,Mpnorm,Pnorm,smanorm,inorm,b)  #b is already roughly in the range (0,1); no need to normalize again\n\n\n  return star_params,star_params_norm\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-09T03:06:16.991819Z","iopub.execute_input":"2025-10-09T03:06:16.992223Z","iopub.status.idle":"2025-10-09T03:06:17.026119Z","shell.execute_reply.started":"2025-10-09T03:06:16.992206Z","shell.execute_reply":"2025-10-09T03:06:17.025415Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# fit functions (for cpu)","metadata":{}},{"cell_type":"code","source":"def hess_estimator(x,*args):\n  data=args[0:-2]\n  chisquare_fn=args[-2]\n  central=args[-1]\n\n  data_for_hess=tuple(list(data)+[False])\n  hess=torch.func.hessian(chisquare_fn)(torch.tensor(central.x,device=cpu_fit_device),*data_for_hess)\n  return hess.detach().cpu().numpy()\n\ncpu_fit_errnames=[\n        \"raw_fit_depth_hess_err\",\n        \"fit_p3_hess_err\",\n        \"fit_p2_hess_err\",\n        \"fit_p1_hess_err\",\n        \"fit_p0_hess_err\",\n        \"Tcenter_hess_err\",\n        \"T_hess_err\",\n        \"tau_hess_err\",\n        \"ldc_hess_err\"\n      ]\ndef check_bounds_cpu(hess_uncerts):\n  if nn_input_var_min is None or nn_input_var_max is None:\n    print(\"check_bounds_cpu called with no bounds\",flush=True)\n    return False\n  if len(hess_uncerts)==0:\n    #print(\"check_bounds_cpu called with empty hess_uncerts list\",flush=True)\n    return False\n  for iname,name in enumerate(cpu_fit_errnames):\n    if name not in nn_input_var_min or name not in nn_input_var_max:\n      print(str(name)+\" appears to be missing from the nn_input_var bounds\",flush=True)\n      return False\n    val=hess_uncerts[iname]\n    if val<nn_input_var_min[name] or val>nn_input_var_max[name]:\n      return False\n  return True\n    \ndef hess_errors_cpu(data,chisquare_fn,central):\n  data_for_hess=tuple(list(data)+[False])\n  hess=torch.func.hessian(chisquare_fn)(torch.tensor(central.x,device=cpu_fit_device),*data_for_hess)\n  #print(\"hess.shape=\"+str(hess.shape),flush=True)\n    \n  try:\n    hessinv=torch.linalg.inv(hess+1e-8*torch.eye(hess.shape[0]).to(cpu_fit_device))\n    errmat=hessinv\n    #print(\"errmat.shape=\"+str(errmat.shape),flush=True)\n    if any([errmat[i,i].item()<0 for i in range(hessinv.shape[0])]):\n      #print(\"negative diagonal elements: \"+str([errmat[i,i] for i in range(hessinv.shape[0])]),flush=True)\n      raise Exception(\"inverse hessian has negative elements on its diagonal\")\n    hess_uncerts=[errmat[i,i].item()**0.5 for i in range(hessinv.shape[0])]\n  except Exception as e:\n    print(\"hessian is not invertible; may need to redo minimization.  Error message: \"+str(e),flush=True)\n    hess_uncerts=[]\n  return hess_uncerts\n\ndef cpu_fit(arg_tup,is_retry=False): \n  chisquare_fn,init_data,data,kwargs_dict=arg_tup\n  central=None\n  nparam=4\n  bounds=None\n  eps=0.00001\n  force_refit=False\n\n  if \"central\" in kwargs_dict:\n    central=kwargs_dict[\"central\"]\n  if \"nparam\" in kwargs_dict:\n    nparam=kwargs_dict[\"nparam\"]\n  if \"bounds\" in kwargs_dict:\n    bounds=kwargs_dict[\"bounds\"]\n  if \"eps\" in kwargs_dict:\n    eps=kwargs_dict[\"eps\"]\n  if \"force_refit\" in kwargs_dict:\n    force_refit=kwargs_dict[\"force_refit\"]\n\n  if central is None or force_refit:\n    #print(\"about to call bfgs with init_data=\"+str(init_data),flush=True)\n    #print(\"...and data=\"+str(data),flush=True)\n    #print(\"dtypes for tensors in init_data=\"+str([x.dtype if torch.is_tensor(x) else None for x in init_data]),flush=True)\n    #print(\"dtypes for tensors in data=\"+str([x.dtype if torch.is_tensor(x) else None for x in data]),flush=True)\n\n      \n    #print(\"dtypes for numpy arrays in init_data=\"+str([x.dtype if isinstance(x,np.ndarray) else None for x in init_data]),flush=True)\n    #print(\"dtypes for numpy arrays in data=\"+str([x.dtype if isinstance(x,np.ndarray) else None for x in data]),flush=True)\n      \n    #print(\"bounds=\"+str(bounds),flush=True)\n    res=scipy.optimize.minimize(chisquare_fn,init_data,\n                                bounds=bounds,\n                                method=\"Newton-CG\",\n                                args=data,jac=True) #,hess=hess_estimator)\n    hess_uncerts=hess_errors_cpu(data,chisquare_fn,res)\n    print(\"newton-cg returns fun=\"+str(res.fun)+\", success=\"+str(res.success)+\", hess_uncerts=\"+str(hess_uncerts),flush=True)\n\n    #switching SLSQP-->Nelder-Mead here seems to recover nearly-identical fit results for the FGS fit, \n    #at least for the one example I have so far checked in detail.\n    #for minim in [\"SLSQP\",\"Powell\"]:\n    #for minim in [\"Newton-CG\",\"TNC\",\"Powell\",\"Nelder-Mead\"]:\n    for minim in [\"TNC\",\"Powell\",\"Nelder-Mead\"]:\n      #for minim in [\"Newton-CG\",\"Powell\"]:\n\n      #count calls to Nelder-Mead over the run and stop when you hit some threshold\n      global n_neldermead\n      if minim==\"Nelder-Mead\":\n        if n_neldermead>=n_neldermead_max:\n          break\n        else:\n          n_neldermead+=1\n            \n      if not res.success or len(hess_uncerts)==0 or not check_bounds_cpu(hess_uncerts) or res.fun>560:  #560 is roughly chisq/ndof=1.5\n        print(\"retry minimization with \"+str(minim),flush=True)\n        #print(\"...check_bounds_cpu returned \"+str(check_bounds_cpu(hess_uncerts)),flush=True)\n        #print(\"starting new call to minimize, method=\"+str(minim)+\"; res.x=\"+str(res.x),flush=True)\n        #if torch.is_tensor(res.x) or isinstance(res.x,np.ndarray):\n        #  print(\"res.x is a tensor or ndarray with dtype=\"+str(res.x.dtype),flush=True)\n        #elif isinstance(res.x,list):\n        #  print(\"res.x is a list whose element types are \"+str([(type(x),x.dtype) if torch.is_tensor(x) or isinstance(x,np.ndarray) else type(x) for x in res.x]),flush=True)\n\n        res=scipy.optimize.minimize(chisquare_fn,res.x,\n                                    bounds=bounds,\n                                    method=minim,\n                                    args=data,jac=True)  #,hess=hess_estimator)\n        hess_uncerts=hess_errors_cpu(data,chisquare_fn,res)\n        print(str(minim)+\" returns fun=\"+str(res.fun)+\", success=\"+str(res.success)+\", hess_uncerts=\"+str(hess_uncerts),flush=True)\n\n    #if len(hess_uncerts)>0:\n    #  print(\"hess_errors_cpu has returned a non-empty list with \"+str(sum([1 if e is None else 0 for e in hess_uncerts]))+\" None(s)\",flush=True)\n    #  print(\"contents: \"+str(hess_uncerts),flush=True)\n    #else:\n    #  print(\"hess_errors_cpu has returned an empty list\",flush=True)\n        \n    if central is None:\n      central=res\n    elif res.fun<central.fun:\n      central=res\n\n  return central,hess_uncerts\n\ndef compute_polynomial_fit(params,obs,include_signal=True):\n  #parameter list, with usual default values:\n  #  - signal norm: [0.]\n  #  - background shape: [0]*NPARAM-1 + [1.]\n  #  - signal center logit: [Tcenter_guess]  --> computed as time_bins*torch.sigmoid(param)\n  #  - transit duration: [T_guess]\n  #  - ingress/egress duration logit: [tau_guess] --> gets computed as (T/2)*torch.sigmoid(param)\n  #  - limb darkening coeff logits: [0.,0.] --> actual coeffs are computed as MAX_LDC_COEF*torch.sigmoid(param)\n\n  obs_x,Tcenter_guess,T_guess=obs\n \n  params=[torch.tensor(c).to(cpu_fit_device) if not torch.is_tensor(c) else c.to(cpu_fit_device) for c in params]\n  if not torch.is_tensor(obs_x):\n    obs_x=torch.tensor(obs_x,device=cpu_fit_device)\n\n  offset=torch.nn.functional.softplus(params[0],beta=100)\n  coeffs=params[1:-4]\n  coeffs=coeffs[::-1]   \n  #coeffs[1:]=[c/100 for c in coeffs[1:]]\n  coeffs[1:]=[0.01*torch.tanh(c) for c in coeffs[1:]]\n\n  fit_y=torch.zeros_like(obs_x,device=cpu_fit_device)\n  for i in range(len(coeffs)):\n    fit_y+=coeffs[i]*obs_x**i\n\n  if include_signal:\n      \n    Tcenter=torch.clip(Tcenter_guess+(obs_x.shape[-1]/2)*torch.tanh(torch.clip(params[-4],-4,4)),10,obs_x.shape[-1]-10)\n    T=torch.clip(T_guess+(obs_x.shape[-1]/2)*torch.tanh(torch.clip(params[-3],-4,4)),1,obs_x.shape[-1])\n    tau=(T/2)*torch.sigmoid(torch.clip(params[-2],-5,5))\n    ldc1=MAX_LDC_COEF*torch.sigmoid(torch.clip(params[-1],-5,5))\n\n    tstart=Tcenter-T/2  \n    tend=Tcenter+T/2\n    ingr_start=tstart-tau/2\n    ingr_end=tstart+tau/2\n    egr_start=tend-tau/2\n    egr_end=tend+tau/2\n\n    tstart=torch.clip(tstart,0,fit_y.shape[-1]-1)\n    tend=torch.clip(tend,0,fit_y.shape[-1]-1)\n    ingr_start=torch.clip(ingr_start,0,fit_y.shape[-1]-1)\n    ingr_end=torch.clip(ingr_end,0,fit_y.shape[-1]-1)\n    egr_start=torch.clip(egr_start,0,fit_y.shape[-1]-1)\n    egr_end=torch.clip(egr_end,0,fit_y.shape[-1]-1)\n\n    in_transit=torch.clip((torch.arange(0,obs_x.shape[-1]).to(obs_x.device)-tstart)/tau,0,1)*torch.clip((tend-torch.arange(0,obs_x.shape[-1]).to(obs_x.device))/tau,0,1)\n    in_transit=in_transit/torch.max(in_transit,dim=-1,keepdim=True).values\n    limbdark=torch.ones_like(obs_x,device=cpu_fit_device)\n    limbdark_x=(torch.arange(0,obs_x.shape[-1],device=cpu_fit_device)-Tcenter)/max(T/2,1)\n    limbdark=torch.clip(limbdark-ldc1*limbdark_x**2,0,1)  \n    limbdark=limbdark/torch.max(limbdark)\n    in_transit=in_transit*limbdark\n    in_transit=in_transit/torch.max(in_transit,dim=-1,keepdim=True).values\n\n    fit_y=fit_y*((1-in_transit*offset))   #torch.where(in_transit>=1,offset,0)\n\n    #print(\"compute_poly: params=\"+str([p.item() for p in params]),flush=True)\n    #print(\"Tcenter,T,tau,ldc1=\"+str([Tcenter.item(),T.item(),tau.item(),ldc1.item()]),flush=True)\n    #print(\"sum(in_transit)=\"+str(torch.sum(in_transit).item()),flush=True)\n    #raise Exception(\"Stop\")\n  return fit_y\n\ndef polynomial_penalty(params):\n  if not torch.is_tensor(params):\n    return 0\n  retval= max(0,torch.abs(params[-4])-4)**2+ max(0,torch.abs(params[-3])-4)**2+  max(0,torch.abs(params[-2])-4)**2+ max(0,torch.abs(params[-1])-4)**2\n  return retval*100\n\ndef light_curve_chisquare(params,*args):\n  #reminder: ch0_signal.shape=torch.Size([1, 288, 5625])\n  #but we are fitting one wavelength at a time so that the optimizer \n  #doesn't get screwed up exploring stupid cross-wavelength correlations.\n  #so we expect arguments to have the following shapes:\n  #  - obs_x, obs_y, obs_yerr, and in_transit:  (5625,)\n  #  - offset and polynomial coefficients: scalar\n        \n  obs_x=args[0]  #torch.tensor(args[0])\n  obs_y=torch.tensor(args[1],device=cpu_fit_device)\n  obs_yerr=torch.tensor(args[2],device=cpu_fit_device)                    \n  compute_func=args[3]\n  penalty_eval_func=args[4]\n  do_grad=True\n  if len(args)==6:\n    do_grad=args[5] \n  \n  if not torch.is_tensor(params):\n    params=torch.tensor(params,device=cpu_fit_device)\n  assert len(params.shape)==1\n  params.requires_grad=True\n\n  fit_y=compute_func(params,obs_x,include_signal=True)\n\n  chisq=(torch.abs(obs_y-fit_y))/obs_yerr  \n  chisq=chisq**2\n  chisq=torch.sum(chisq)\n\n  if penalty_eval_func is not None:\n    #print(\"chisq mean=\"+str(chisq.mean().item())+\" before penalty....\",flush=True)\n    chisq=chisq+penalty_eval_func(params)\n    #print(\"...and \"+str(chisq.mean().item())+\" after\",flush=True)\n\n  if not do_grad:\n    return chisq\n\n  chisq.backward()\n  if torch.any(torch.abs(params.grad)>1000):\n    scalefac=torch.max(torch.abs(params.grad)).detach() \n    params.grad=params.grad/scalefac\n\n  grad=params.grad.detach().cpu().numpy()\n  #print(\"grad=\"+str(grad),flush=True)\n  #raise Exception(\"stop\")\n    \n  return chisq.detach().cpu().numpy(),grad\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-09T03:06:17.02703Z","iopub.execute_input":"2025-10-09T03:06:17.027259Z","iopub.status.idle":"2025-10-09T03:06:17.054093Z","shell.execute_reply.started":"2025-10-09T03:06:17.027238Z","shell.execute_reply":"2025-10-09T03:06:17.053354Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# fit functions (for gpu)\n","metadata":{}},{"cell_type":"code","source":"ChisqFitResult=collections.namedtuple(\"ChisqFitResult\",[\"success\",\"fun\",\"x\",\"instance\"])\n\nclass LightCurveChisquare(torch.nn.Module):\n  def __init__(self,init_data): \n    super().__init__()\n    device=init_data[0].device\n    self.device=device\n    self.cached_signal_shape=None\n\n    #assert len(init_data[0].shape)==1\n    self.n_wavelengths=init_data[0].shape[0]\n\n    self.signal_offset=torch.nn.parameter.Parameter(data=init_data[0].to(device))\n    idx=0\n    self.tau_param=torch.nn.parameter.Parameter(data=init_data[1].to(device))\n    self.ldc_param=torch.nn.parameter.Parameter(data=init_data[2].to(device))\n    \n    #polynomial parameters\n    self.poly_c0=torch.nn.parameter.Parameter(data=init_data[3].to(device))\n    self.poly_c1=torch.nn.parameter.Parameter(data=init_data[4].to(device))\n    self.poly_c2=torch.nn.parameter.Parameter(data=init_data[5].to(device))\n    self.poly_c3=torch.nn.parameter.Parameter(data=init_data[6].to(device))\n\n  def forward(self,data,fixpars):\n    params=[self.signal_offset,self.tau_param,self.ldc_param]\n    params+=[self.poly_c0,self.poly_c1,self.poly_c2,self.poly_c3]\n\n    #these penalty terms keep the tau and ldc parameters from wandering off far beyond the range \n    #over which we have computed interpolation templates\n    penalty=torch.clip(torch.abs(self.tau_param)-5,0,1e30)\n    penalty+=torch.clip(torch.abs(self.ldc_param)-5,0,1e30)\n\n    return self.compute_chisq(params,data)+torch.sum(penalty)\n\n  def compute_chisq(self,params,data):\n    obs_x=data[0].to(self.device)\n    obs_y=data[1].to(self.device)\n    obs_yerr=data[2].to(self.device)\n    Tcenter_guess,T_guess,tau_guess,MAX_LDC_COEF,prefit_x=data[3]\n\n    batch_shape=obs_y.shape[:-1]\n    for i in range(len(obs_y.shape[:-1])):\n      obs_x=torch.unsqueeze(obs_x,dim=0)\n\n    if not torch.is_tensor(obs_x):\n      obs_x=torch.tensor(obs_x).to(self.device)\n\n\n    poly_y,signal=gpu_compute_poly_fit_func(params,(obs_x,Tcenter_guess,T_guess,tau_guess,MAX_LDC_COEF,prefit_x),include_signal=True)\n\n    chisq=(torch.abs(obs_y-poly_y)/torch.clip(obs_yerr,min=1e-4))**2\n    chisq=torch.sum(chisq,dim=-1)\n    return chisq\n\n  #convenience function to extract fit results from the chisquare model above and expose them via an interface that looks a bit like scipy.optimize.OptimizeResult\n  def fit_result(self,success,loss):\n    params=[self.signal_offset,self.tau_param,self.ldc_param,self.poly_c0,self.poly_c1,self.poly_c2,self.poly_c3]\n\n    params=[p.data.detach().cpu() for p in params]\n    return ChisqFitResult(success=True,fun=loss,x=params,instance=self)\n\n#@torch.compile\ndef gpu_compute_poly_fit_func(params,obs,include_signal=True,fixpars=None,return_signal=False,signal_shape=None):\n    obs_x,Tcenter_guess,T_guess,tau_guess,MAX_LDC_COEF,prefit_x=obs\n\n    #reminder: input params=[self.signal_offset,self.Tcenter_param,self.T_param,self.tau_param,self.ldc_param,self.poly_c0,self.poly_c1,self.poly_c2,self.poly_c3]\n    #if not free_signal_shape, then self.Tcenter_param, self.T_param, self.tau_param, and self.ldc_param are absent\n    offset=torch.nn.functional.softplus(params[0],beta=100)\n\n\n    #Tcenter_param=torch.clip(Tcenter_guess+(obs_x.shape[-1]/2)*prefit_x[-4],0,obs_x.shape[-1])\n    #T_param=torch.clip(T_guess+(obs_x.shape[-1]/2)*torch.nn.functional.softplus(prefit_x[-3],beta=100),0,obs_x.shape[-1])\n\n\n    Tcenter_param=prefit_x[0].to(offset.device)*torch.ones_like(offset,device=offset.device)\n    T_param=prefit_x[1].to(offset.device)*torch.ones_like(offset,device=offset.device)\n    tau_param=params[1]\n    ldc_param=params[2]\n    coeffs=[p.to(offset.device) for p in params[3:]]\n    coeffs[1:]=[0.01*torch.tanh(c) for c in coeffs[1:]]\n\n    if len(offset.shape)==0:\n      offset=torch.unsqueeze(offset,dim=0)\n      Tcenter_param=torch.unsqueeze(Tcenter_param,dim=0)\n      T_param=torch.unsqueeze(T_param,dim=0)\n      tau_param=torch.unsqueeze(tau_param,dim=0)\n      ldc_param=torch.unsqueeze(ldc_param,dim=0)\n      coeffs=[torch.unsqueeze(c,dim=0) for c in coeffs]\n\n    fit_y=torch.unsqueeze(coeffs[0],dim=-1)\n    for ic,c in enumerate(coeffs[1:]):\n      cc=torch.unsqueeze(c,dim=-1)\n      ox=obs_x\n      fit_y=fit_y+cc*ox**(ic+1)\n\n    in_transit=None\n    if include_signal:\n      if len(offset.shape)==1:\n        off=torch.unsqueeze(offset,dim=-1)\n      else:\n        off=offset\n\n      Tcenter=torch.clip(Tcenter_guess+(obs_x.shape[-1]/2)*torch.tanh(torch.clip(Tcenter_param,-4,4)),10,obs_x.shape[-1]-10)\n      T=torch.clip(T_guess+(obs_x.shape[-1]/2)*torch.tanh(torch.clip(T_param,-4,4)),0,obs_x.shape[-1])\n      tau=(T/2)*torch.sigmoid(torch.clip(tau_param,-5,5))\n      ldc1=MAX_LDC_COEF*torch.sigmoid(torch.clip(ldc_param,-5,5))\n\n      Tcenter=torch.unsqueeze(Tcenter,dim=-1)\n      T=torch.unsqueeze(T,dim=-1)\n      tau=torch.unsqueeze(tau,dim=-1)\n      ldc1=torch.unsqueeze(ldc1,dim=-1)\n\n      tstart=Tcenter-T/2 \n      tend=Tcenter+T/2\n\n      in_transit=torch.clip((torch.arange(0,obs_x.shape[-1]).to(obs_x.device)-tstart)/tau,0,1)*torch.clip((tend-torch.arange(0,obs_x.shape[-1]).to(obs_x.device))/tau,0,1)\n      in_transit=in_transit/torch.max(in_transit,dim=-1,keepdim=True).values\n\n      limbdark_x=(obs_x*obs_x.shape[-1]-Tcenter)/torch.maximum(T/2,torch.ones_like(T))\n      limbdark=torch.ones_like(obs_x,device=offset.device)\n      limbdark=torch.clip(limbdark-ldc1*limbdark_x**2,0,1)          \n      limbdark=limbdark/torch.max(limbdark,dim=-1,keepdim=True).values\n      in_transit=in_transit*limbdark\n      in_transit=in_transit/torch.max(in_transit,dim=-1,keepdim=True).values\n\n      fit_y=fit_y*(1-in_transit*off)\n\n    return fit_y,in_transit\n\n\n\ndef torch_minimize(chisq_class,init_data,data,fixpars=None,niter=5000,lr=3e-4,lr_decay=0.999,noisy=False,limit=0.01):\n  #chisq_class:  a torch.nn.Module (class, not instance) that computes a chisquare to be minimized.  Its forward() should take two arguments, data and fixpars.\n  #              It should also implement a method, fit_result(success,loss) which returns a ChisqFitResult summarizing the success, loss and best-fit parameters.\n  #data:  the argument to the chisquare model\n  #fixpars:  if not None, a list of 2-tuples.  First element of each tuple is an index on params, indicating which param is to\n  #          be held constant.  Second element is the value at which that param should be held fixed.\n  #          If None, then loss is optimized over all parameters.\n  #\n  loss_fn=chisq_class(init_data)\n  optimizer=torch.optim.Adam(loss_fn.parameters(),lr=lr)\n  scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma=lr_decay)\n\n  last_check=None\n  for i_iter in range(niter):\n    optimizer.zero_grad()\n    loss=loss_fn(data,fixpars)\n    loss_sum=torch.sum(loss)\n\n    loss_sum.backward()\n    optimizer.step()\n    scheduler.step()\n\n    if i_iter>niter/8 and i_iter%10==0:\n      loss_mean=loss.mean().item()\n      if last_check is None:\n        last_check=loss_mean\n      elif last_check-loss_mean<limit and torch.max(loss)<560: \n        print(\"minimization terminates at i_iter=\"+str(i_iter)+\" because loss is not dropping very much anymore\",flush=True)\n        break\n      elif loss_mean<last_check:\n        last_check=loss_mean\n  else:\n    print(\"optimization loop ran all the way to iteration \"+str(i_iter),flush=True)\n\n  #Optimized parameters are stored in the model; query that to get the fit result.\n  res=loss_fn.fit_result(True,loss.detach().cpu())\n  if fixpars is not None:\n    for iparam,p in fixpars:\n      res.x[iparam]=p.detach()  #.cpu()\n  return res\n\n\ndef compute_hess_uncerts(central,chisq_class,init_data,data,device,eps=1e-8):\n  hess_uncerts=[]\n  last=None\n  for ibatch in range(central.x[0].shape[-1]):\n    if last is not None:\n      hess_uncerts.append(last)\n      last=None\n      continue\n        \n    chisq_class_instance=chisq_class([torch.unsqueeze(p[ibatch],dim=0).to(hess_device) for p in central.x])\n    data_for_hess=tuple([data[0].to(hess_device),torch.unsqueeze(data[1][ibatch,:],dim=0).to(hess_device),torch.unsqueeze(data[2][ibatch,:],dim=0).to(hess_device),data[3]])\n    hparam=torch.cat([torch.unsqueeze(p[ibatch],dim=0).to(hess_device) for p in central.x],dim=0)\n    hess=torch.func.hessian(chisq_class_instance.compute_chisq)(hparam,data_for_hess)\n    hess=torch.squeeze(hess,dim=0)\n      \n    try:\n      hessinv=torch.linalg.inv(hess+eps*torch.eye(hess.shape[0]).to(hess_device))\n      errmat=hessinv\n      if any([errmat[i,i].item()<0 for i in range(hessinv.shape[-1])]):\n        hess_uncerts.append(None)\n      else:\n        hess_uncerts.append(torch.tensor([[errmat[i,i].item()**0.5 for i in range(hessinv.shape[-1])]],device=torch.device('cpu')))\n        last=hess_uncerts[-1]\n    except Exception as e:\n      hess_uncerts.append(None)\n\n  return hess_uncerts\n\n\ngpu_fit_errnames=[\n        \"raw_fit_depth_hess_err\",\n        \"tau_hess_err\", \n        \"ldc_hess_err\",\n        \"fit_p0_hess_err\",\n        \"fit_p1_hess_err\",\n        \"fit_p2_hess_err\",\n        \"fit_p3_hess_err\"\n]\n\ndef check_bounds_gpu(hess_uncerts):\n  if nn_input_var_min is None or nn_input_var_max is None:\n    #print(\"check_bounds_cpu called with no bounds\",flush=True)\n    return True\n  #else:\n  #  print(\"check_bounds_gpu called with hess_uncerts.shape=\"+str(hess_uncerts.shape),flush=True)\n  for iname,name in enumerate(gpu_fit_errnames):\n    if name not in nn_input_var_min or name not in nn_input_var_max:\n      print(str(name)+\" appears to be missing from the nn_input_var bounds\",flush=True)\n      return True\n    val=hess_uncerts[0,iname]\n    #if name in [\"fit_p1_hess_err\",\"fit_p2_hess_err\",\"fit_p3_hess_err\"]:\n    #    val=val*100  #because the bounds file is based on rescaled values for these params, and rescaling has not happened yet\n    if val<nn_input_var_min[name] or val>nn_input_var_max[name]:\n      print(\"check_bounds_gpu returns false for \"+str(name)+\": val=\"+str(val)+\" but bounds=\"+str([nn_input_var_min[name],nn_input_var_max[name]]),flush=True)\n      return False\n  return True\n    \ndef gpu_fit(chisq_class,init_data,data,device,central=None,niter=5000,niter_errfit=1000,lr=3e-4,lr_decay=0.999,eps=0.001,noisy=False,force_refit=False,depth=0):\n\n  if central is None or force_refit:\n    new_central=torch_minimize(chisq_class,init_data,data,niter=niter,lr=lr,lr_decay=lr_decay,noisy=noisy,limit=1.0)\n    print(\"at depth=\"+str(depth)+\", new_central mean function value=\"+str(torch.mean(new_central.fun).item()),flush=True)\n    if central is None:\n      central=new_central\n    else:\n      x=[torch.where(new_central.fun.to(device)<central.fun.to(device),new_central.x[i].to(device),central.x[i].to(device)).detach() for i in range(len(central.x))]\n      new_init_data=x\n      fun=torch.minimum(new_central.fun.to(device),central.fun.to(device))\n      central=ChisqFitResult(success=False,fun=fun,x=x,instance=chisq_class(new_init_data))\n\n  print(\"finished central fit; working on hessian uncertainties...\",flush=True)\n\n  hess_uncerts=compute_hess_uncerts(central,chisq_class,init_data,data,device)\n\n  if all([u is None for u in hess_uncerts]) or len(hess_uncerts)==0:\n    print(\"empty hess_uncerts in gpu_fit\",flush=True)\n    hess_uncerts=[]\n    errnames=[\"raw_fit_depth_hess_err\",\"tau_hess_err\",\"ldc_hess_err\",\"fit_p0_hess_err\",\"fit_p1_hess_err\",\"fit_p2_hess_err\",\"fit_p3_hess_err\"]\n    for j in range(len(avg_hess_uncerts)-1):\n      default=[avg_hess_uncerts.iloc[j+1][name] for name in errnames]\n      #the default values have the scale factor of 100 applied to p1/p2/p3 and their uncertainties, \n      #but the values normally coming out of this function do not. \n      # --> no longer the case in v19-debugNCG\n      #default[4]/=100\n      #default[5]/=100\n      #default[6]/=100\n      hess_uncerts.append(torch.tensor([default]))\n    print(\"done dealing with that\") \n  #else:\n  #  #crazy pills: check uncertainties against the ranges the neural nets were trained on, and re-minimize if any are out of range\n  #  crazy_pills=False\n  #  for iu,u in enumerate(hess_uncerts):\n  #    if u is None:\n  #      continue\n  #    cb=check_bounds_gpu(u)\n  #    if not cb:\n  #        print(\"crazy-pills check fails for wavelength \"+str(iu),flush=True)\n  #    crazy_pills=crazy_pills or not cb\n  #  if crazy_pills:\n  #    print(\"some fit uncertainties are outside the expected range; try re-minimizing...\",flush=True)\n  #    new_central=torch_minimize(chisq_class,central.x,data,niter=niter,lr=lr,lr_decay=lr_decay,noisy=noisy,limit=0.01) #1.0)\n  #    print(\"after crazy-pills retry, new_central mean function value=\"+str(torch.mean(new_central.fun).item()),flush=True)\n\n  #    x=[torch.where(new_central.fun.to(device)<central.fun.to(device),new_central.x[i].to(device),central.x[i].to(device)).detach() for i in range(len(central.x))]\n  #    new_init_data=x\n  #    fun=torch.minimum(new_central.fun.to(device),central.fun.to(device))\n  #    central=ChisqFitResult(success=False,fun=fun,x=x,instance=chisq_class(new_init_data))\n  #    hess_uncerts=compute_hess_uncerts(central,chisq_class,init_data,data,device)\n\n  #    crazy_pills=False\n  #    for iu,u in enumerate(hess_uncerts):\n  #      if u is None:\n  #        continue\n  #      cb=check_bounds_gpu(u)\n  #      if not cb:\n  #        print(\"crazy-pills re-check fails for wavelength \"+str(iu),flush=True)\n  #      crazy_pills=crazy_pills or not cb\n  #    print(\"after re-minimizing, check on uncertainty bounds returns: \"+str(crazy_pills),flush=True)\n        \n        \n  #print(\"gpu compute_hess_uncerts has returned a list with \"+str(sum([1 if e is None else 0 for e in hess_uncerts]))+\" None values\",flush=True)\n  #fill in any missing uncertainties with extrapolated/interpolated values\n  for i in range(len(hess_uncerts)):\n    if hess_uncerts[i] is None:\n      if i==0 or all([hess_uncerts[j] is None for j in range(i)]):\n        first_not_none=min([j for j in range(len(hess_uncerts)) if hess_uncerts[j] is not None])\n        hess_uncerts[i]=hess_uncerts[first_not_none]\n      elif i==len(hess_uncerts)-1 or all([hess_uncerts[j] is None for j in range(i,len(hess_uncerts))]):\n        last_not_none=max([j for j in range(len(hess_uncerts)) if hess_uncerts[j] is not None])\n        hess_uncerts[i]=hess_uncerts[last_not_none]\n      else:\n        prev=max([j for j in range(0,i) if hess_uncerts[j] is not None])\n        nxt=min([j for j in range(i,len(hess_uncerts)) if hess_uncerts[j] is not None])\n        hess_uncerts[i]=(hess_uncerts[prev]+hess_uncerts[nxt])/2\n\n  hess_uncerts=torch.cat(hess_uncerts,dim=0)\n\n  #print(\"returning from gpu_fit, hess_uncerts=\"+str(hess_uncerts),flush=True)\n  return len(hess_uncerts)>0,central,hess_uncerts\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-09T03:06:17.054906Z","iopub.execute_input":"2025-10-09T03:06:17.055393Z","iopub.status.idle":"2025-10-09T03:06:17.091131Z","shell.execute_reply.started":"2025-10-09T03:06:17.055368Z","shell.execute_reply":"2025-10-09T03:06:17.090378Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# functions for bias correction neural net","metadata":{}},{"cell_type":"code","source":"#neural net for limb darkening correction, bias correction, and smoothing.\n#Originally intended to also interpolate binned fit results to get predictions for each wavelength.\n#Wound up not doing the binning, but the name stuck.\nclass CorrInterp(torch.nn.Module):\n  def __init__(self,block_dict,n_hidden,n_hidden_pe,rebinnings,device,ref):\n    super().__init__()\n    self.n_hidden=n_hidden\n    self.n_hidden_pe=n_hidden_pe\n    self.rebinnings=rebinnings\n    self.device=device\n    self.ref=ref\n\n    if ref==\"fit\":\n      self.nfeats_in=sum([block_dict[key].shape[-1] for key in [\"baseline\",\"poly\",\"res\",\"sig\",\"bg\",\"spatial\"]])\n      self.nfeats_in+=2*(len(rebinnings)+1)*(block_dict[\"sliding_window_in\"].shape[-1])  #for the \"sliding_window_<size>\" and \"sliding_window_err_<size>\" blocks\n    else:\n      self.nfeats_in=sum([block_dict[key].shape[-1] for key in [\"baseline\",\"spatial\"]])\n\n    self.fex=torch.nn.Sequential(OrderedDict([\n      (\"fc1\",torch.nn.Linear(self.nfeats_in,2*self.nfeats_in,device=device)),\n      (\"relu\",torch.nn.LeakyReLU()),\n      (\"dropout\",torch.nn.Dropout()),\n      (\"fc2\",torch.nn.Linear(2*self.nfeats_in,n_hidden,device=device))\n    ]))\n\n    #self.proj=torch.nn.Linear(n_hidden,n_hidden_pe,device=device)\n\n    self.conv1=torch.nn.Conv1d(n_hidden_pe,n_hidden_pe,3,padding=1,device=device)\n    self.relu=torch.nn.LeakyReLU()\n    self.conv3=torch.nn.Conv1d(n_hidden_pe,2,3,padding=1,device=device)  #there was a conv2 in some variants of this class\n\n\n  def forward(self,data,mask=None,width_scale=None):\n    #input blocks should be shape (batch_size,n_wavelength_bins,feat_size)\n    #fex will output shape (batch_size,n_wavelength_bins,n_hidden)\n    #Need to transpose the last two to make inputs for conv layers \n    #feats=[torch.transpose(fex(data[var]),-1,-2) for var,fex in self.fex_dict.items()]\n    #print(\"welcome to CorrInterp::forward\",flush=True)\n\n    #during training, checked that this was the case for key in [0,1,2]\n    #assert \"baseline\" in data[key]\n    #assert \"fit\" in data[key]\n    #assert \"poly\" in data[key]\n    #assert \"res\" in data[key]\n    #assert \"bg\" in data[key]\n    #assert \"sig\" in data[key]\n    #assert \"baseline_unnorm\" in data[key]\n    #assert \"fit_unnorm\" in data[key]\n\n    if self.ref==\"fit\":\n      pred=data[\"fit_unnorm\"][...,0:2].clone()  #raw fit mean and uncert\n      block_list=[\"baseline\",\"poly\",\"res\",\"sig\",\"bg\",\"spatial\"]\n    else:\n      pred=data[\"baseline_unnorm\"][...,0:2].clone()\n      block_list=[\"baseline\",\"spatial\"]\n\n    feats=[]\n    for key in block_list:\n      #print(\"key=\"+str(key)+\", feats list gets \"+str(data[0][key]),flush=True)\n      feats.append(data[key])\n\n    #for key in [\"sliding_window_in\",\"sliding_window_errs_in\"]+[\"sliding_window_\"+str(sw_size) for sw_size in self.rebinnings]+[\"sliding_window_errs_\"+str(sw_size) for sw_size in self.rebinnings]:\n    #  print(\"key=\"+str(key)+\", feats list gets \"+str(data[0][key]),flush=True)\n    if self.ref==\"fit\":      \n      feats.append(data[\"sliding_window_in\"])\n      feats.append(data[\"sliding_window_errs_in\"])\n      for sw_size in self.rebinnings:\n        feats.append(data[\"sliding_window_\"+str(sw_size)])\n        feats.append(data[\"sliding_window_errs_\"+str(sw_size)])\n    #print(\"len(feats)=\"+str(len(feats)),flush=True)\n\n    #print(\"feats dtypes: \"+str([x.dtype for x in feats]),flush=True)\n    x=torch.cat(feats,dim=-1)\n    #print(\"x.shape=\"+str(x.shape),flush=True)\n    #print(\"fex gets input x=\"+str(x),flush=True)\n\n    #print(\"x.shape=\"+str(x.shape),flush=True)\n    #print(\"x=\"+str(x),flush=True)\n    #print(\"row 0 as list: \"+str(x[0,0,:].tolist()),flush=True)\n    #print(\"row 1 as list: \"+str(x[0,1,:].tolist()),flush=True)\n\n    feats=self.fex(x)\n    feats=torch.transpose(feats,1,2)  #(batch_size,n_hidden,n_wavelength_bins)\n\n    x=self.relu(self.conv1(feats))\n    x=torch.nn.functional.dropout(x,training=self.training)\n    res=self.conv3(x)  #shape (batch_size,2*len(rebinnings)+1,n_wavelength_bins)\n\n    #initial training is smoother if you have a bigger uncertainty\n    if width_scale is not None:\n      pred[...,1]=pred[...,1]*width_scale\n\n    #inverse softplus to get a logit (which can be negative) instead of a sigma (which can't)\n    pred[...,1]=torch.log(torch.expm1(torch.clip(pred[...,1],1e-6,1.0)))\n    if torch.any(torch.isnan(pred)):\n      print(\"nan in fit_unnorm after inverse softplus in model forward\",flush=True)\n      raise Exception(\"nan in fit_unnorm after inverse softplus in model forward\")\n    if torch.any(torch.isinf(pred)):\n      print(\"inf in fit_unnorm after inverse softplus in model forward\",flush=True)\n      raise Exception(\"inf in fit_unnorm after inverse softplus in model forward\")\n\n    pred=torch.transpose(pred,-1,-2)\n\n    #print(\"forward for ref=\"+str(self.ref)+\" sees pred=\"+str(pred),flush=True)\n    #print(\"...and res=\"+str(res),flush=True)\n      \n    retval=pred+res\n    return retval  #,torch.mean(res[...,1]**2)\n\ndef make_feature(row,var_max=None,var_min=None,orbit_block=None):  #,signal_shape_params=None):\n\n\n  varnames=[(\"baseline\",[\"baseline_depth\",\"baseline_depth_err\",\"norm\"])]\n  varnames.append((\"fit\",[\"raw_fit_depth\",\"raw_fit_depth_hess_err\",\"lincorr_fit_depth\",\"lincorr_fit_depth_err\",\"chisq_ndof\"]))\n  varnames.append((\"poly\",[\"fit_p0\",\"fit_p1\",\"fit_p2\",\"fit_p3\",\"fit_p0_hess_err\",\"fit_p1_hess_err\",\"fit_p2_hess_err\",\"fit_p3_hess_err\"]))\n  varnames.append((\"res\",[\"res_pre\",\"res_ingr\",\"res_mid\",\"res_egr\",\"res_post\"]))\n  varnames.append((\"bg\",[\"bg_pre\",\"bg_ingr\",\"bg_mid\",\"bg_egr\",\"bg_post\"]))\n  varnames.append((\"sig\",[\"sig_ingr\",\"sig_mid\",\"sig_egr\",\"T\",\"T_hess_err\"]))\n  varnames.append((\"spatial\",[\"nmask_scaled\",\"spatial_center_mean\",\"spatial_center_std\",\"spatial_width_mean\",\"spatial_width_std\",\"spatial_center_range\",\"spatial_width_range\"]))\n  varnames.append((\"sliding_window_in\",[\"raw_fit_depth\",\"tau\",\"ldc\"]))\n  varnames.append((\"sliding_window_errs_in\",[\"raw_fit_depth_hess_err\",\"tau_hess_err\",\"ldc_hess_err\"]))\n  varnames.append((\"extras\",[\"tstart_mean\",\"tend_mean\",\"tstart_sig\",\"tend_sig\"]))\n\n\n    \n  #varnames=[(\"baseline\",[\"baseline_depth\",\"baseline_depth_err\",\"norm\"])] \n  #varnames.append((\"fit\",[\"raw_fit_depth\",\"raw_fit_depth_hess_err\",\"lincorr_fit_depth\",\"lincorr_fit_depth_err\",\"chisq_ndof\"]))\n  #varnames.append((\"poly\",[\"fit_p0\",\"fit_p1\",\"fit_p2\",\"fit_p3\",\"fit_p0_hess_err\",\"fit_p1_hess_err\",\"fit_p2_hess_err\",\"fit_p3_hess_err\"]))\n  #varnames.append((\"res\",[\"res_pre\",\"res_ingr\",\"res_mid\",\"res_egr\",\"res_post\"]))\n  #varnames.append((\"bg\",[\"bg_pre\",\"bg_ingr\",\"bg_mid\",\"bg_egr\",\"bg_post\"]))  \n  #varnames.append((\"sig\",[\"sig_ingr\",\"sig_mid\",\"sig_egr\",\"T\",\"T_hess_err\"]))  \n  #varnames.append((\"spatial\",[\"nmask_scaled\",\"spatial_center_mean\",\"spatial_center_std\",\"spatial_width_mean\",\"spatial_width_std\",\"spatial_center_range\",\"spatial_width_range\"]))\n  #varnames.append((\"sliding_window_in\",[\"raw_fit_depth\",\"tau\",\"ldc\"]))\n  #varnames.append((\"sliding_window_errs_in\",[\"raw_fit_depth_hess_err\",\"tau_hess_err\",\"ldc_hess_err\"]))\n  #varnames.append((\"extras\",[\"tstart_mean\",\"tend_mean\",\"tstart_sig\",\"tend_sig\"]))\n\n  #n_yell=0\n  blocks=dict()\n  for key,vlist in varnames:\n    feat=[]\n    feat_unnorm=[]\n    for varname in vlist:\n      feat_unnorm.append(row[varname])\n      if var_max is None or var_min is None or varname not in var_max or varname not in var_min:\n        #print(\"could not find varname=\"+str(varname)+\" in at least one of var_max=\"+str(var_max)+\" or var_min=\"+str(var_min),flush=True)\n        #n_yell+=1\n        #if n_yell>10:\n        #    raise Exception(\"stfu\")   \n        feat.append(row[varname])\n      else:\n        feat.append((row[varname]-var_min[varname])/(var_max[varname]-var_min[varname]))\n        if feat[-1]<0 or feat[-1]>1:\n          #print(\"normalized value of \"+str(varname)+\" in block \"+str(key)+\" is not so normalized!\",flush=True)\n          #print(\"original value: \"+str(row[varname]),flush=True)\n          #print(\"min value: \"+str(var_min[varname]),flush=True)\n          #print(\"max value: \"+str(var_max[varname]),flush=True)\n          #print(\"feat so far: \"+str(feat),flush=True)\n          #print(\"!!!!!------>badly normalized value: \"+str(feat[-1]),flush=True)\n          #if feat[-1]>1.3 or feat[-1]<-0.3:\n          #  raise Exception(\"stupid value: \"+str(feat[-1]))\n              \n          feat[-1]=max(0,min(1,feat[-1]))\n\n      if not torch.is_tensor(feat_unnorm[-1]):\n        feat_unnorm[-1]=torch.tensor(feat_unnorm[-1],dtype=torch.float32,device=device)\n      else:\n        feat_unnorm[-1]=feat_unnorm[-1].to(device)\n      if not torch.is_tensor(feat[-1]):\n        feat[-1]=torch.tensor(feat[-1],dtype=torch.float32,device=device)\n      else:\n        feat[-1]=feat[-1].to(device)\n      if len(feat_unnorm[-1].shape)==0:\n        feat_unnorm[-1]=torch.unsqueeze(feat_unnorm[-1],dim=0)\n      if len(feat[-1].shape)==0:\n        feat[-1]=torch.unsqueeze(feat[-1],dim=0)\n\n    #print(\"for key=\"+str(key)+\", we have the following:\",flush=True)\n    #print(\"  type(feat)=\"+str(type(feat)),flush=True)\n    #print(\"  types _in_ feat: \"+str([type(x) for x in feat]),flush=True)\n    #print(\"  shapes in feat: \"+str([x.shape for x in feat]),flush=True)\n    #print(\"  devices in feat: \"+str([x.device for x in feat]),flush=True)\n      \n    #blocks[key]=np.array([feat],dtype=np.float32)\n    blocks[key]=torch.unsqueeze(torch.cat(feat,dim=-1),dim=0)\n    #print(\"a\",flush=True)\n    #blocks[key+\"_unnorm\"]=np.array([feat_unnorm],dtype=np.float32)\n    blocks[key+\"_unnorm\"]=torch.unsqueeze(torch.cat(feat_unnorm,dim=-1),dim=0)\n    #print(\"b\",flush=True)\n    if torch.any(torch.isnan(blocks[key])):\n      raise Exception(\"block for key=\"+str(key)+\" has nan: \"+str(blocks[key]))\n\n    #print(\"c\",flush=True)\n  #print(\"d\",flush=True)\n  if orbit_block is not None:\n    #print(\"type(orbit_block)=\"+str(type(orbit_block)),flush=True) \n    blocks[\"baseline\"]=torch.cat([blocks[\"baseline\"],orbit_block],axis=-1)\n  if torch.any(torch.isnan(blocks[\"poly\"])) or torch.any(torch.isinf(blocks[\"poly\"])):\n    print(\"blocks[poly]=\"+str(blocks[\"poly\"]),flush=True)\n    raise Exception(\"nan in poly block in make_feature\")\n\n  if torch.any(blocks[\"fit_unnorm\"][...,1]<=0):\n    print(\"bad fit_unnorm block? \"+str(blocks[\"fit_unnorm\"]),flush=True)\n    blocks[\"fit_unnorm\"][...,1]=torch.abs(blocks[\"fit_unnorm\"][...,1])\n\n  return blocks\n\ndef collate_example(blocks):\n  if len(blocks)==0:\n    raise Exception(\"collate_example called with an empty blocks list\")\n      \n  coll=dict()\n  for key in sorted(blocks[0].keys()):\n    #print(\"key: \"+str(key)+\", types=\"+str([type(b[key]) for b in blocks]),flush=True)\n    feat=torch.cat([b[key] for b in blocks],dim=-2)\n    if torch.any(torch.isnan(feat)) or torch.any(torch.isinf(feat)):\n      print(\"key=\"+str(key)+\" feat=\"+str(feat),flush=True)\n      raise Exception(\"nan or inf in feature for key=\"+str(key)+\": \"+str(feat))\n      \n\n    #scan for and remove/impute any uncertainties that are tiny or huge\n    #expect feat has shape (n_wavelengths,n_features_in_block)\n    if key==\"sliding_window_errs_in_unnorm\":\n      mask_small=torch.where(feat<1e-6,0,1)\n      mask_large=torch.where(feat>100,0,1)\n      mask=torch.minimum(mask_small,mask_large).to(dtype=bool)\n      at_least_one_good=torch.any(mask,axis=0)\n      if not torch.all(at_least_one_good):  # and rebin_id==0:\n        print(\"feat=\"+str(feat),flush=True)\n        raise Exception(\"some planet in this batch has a problem with at_least_one_good: \"+str(at_least_one_good))\n\n      denom=torch.sum(mask,dim=0,keepdim=True)\n      means=torch.sum(feat,dim=0,keepdim=True)/denom\n      feat=torch.where(mask,means,feat)\n\n    coll[key]=feat\n\n  #also done here:  compute some derived features by averaging, in a sliding window of various sizes, \n  #                 a few best-fit parameters and their errors \n  for sw_size in REBINNINGS:\n    feat=coll[\"sliding_window_in\"].detach().cpu().numpy()  #expect shape (n_wavelength_bins,n_values_in_block)\n    feat_errs=coll[\"sliding_window_errs_in\"].detach().cpu().numpy()  #same shape\n\n    feat_pad=np.pad(feat,((int((sw_size-1)/2),int((sw_size-1)/2)),(0,0)),mode=\"constant\",constant_values=0.)\n    feat_errs_pad=np.pad(feat_errs,((int((sw_size-1)/2),int((sw_size-1)/2)),(0,0)),mode=\"constant\",constant_values=1e10)\n    sw=np.lib.stride_tricks.sliding_window_view(feat_pad,window_shape=sw_size,axis=0)  #expect shape (n_wavelength_bins,n_values_in_block,sw_size)\n    sw_errs=np.lib.stride_tricks.sliding_window_view(feat_errs_pad,window_shape=sw_size,axis=0)\n    weights=np.where(sw_errs>0,1./sw_errs**2,0.)\n    sumweights=np.sum(weights,axis=-1,keepdims=True)\n    weights=np.where(sumweights>0,weights/sumweights,0)\n    feat=np.sum(sw*weights,axis=-1)\n    feat_errs=np.where(sumweights>0,(1./sumweights**0.5),0)[...,0]\n    mask=np.sum(np.where(np.abs(sw)>0,1,0),axis=-1)\n    feat_std=np.where(mask>1,np.std(sw,axis=-1),0)\n    mask=np.sum(np.where(np.abs(sw_errs)>0,1,0),axis=-1)\n    feat_err_std=np.where(mask>1,np.std(np.where(sw_errs>100,0,sw_errs),axis=-1),0)\n    coll[\"sliding_window_\"+str(sw_size)]=torch.tensor(feat,dtype=torch.float32,device=device)\n    coll[\"sliding_window_errs_\"+str(sw_size)]=torch.tensor(feat_errs,dtype=torch.float32,device=device)\n    coll[\"sliding_window_std_\"+str(sw_size)]=torch.tensor(feat_std,dtype=torch.float32,device=device)\n    coll[\"sliding_window_errs_std_\"+str(sw_size)]=torch.tensor(feat_err_std,dtype=torch.float32,device=device)\n\n  #print(\"starting last step of collate\",flush=True)  \n  #finally, give every block a batch index; even though we will likely run examples through the postproc\n  #network one at a time, that network was set up to expect a batch index. \n  for key in coll.keys():\n    coll[key]=torch.unsqueeze(coll[key],dim=0).to(torch.float32)\n  #print(\"all done with collate\",flush=True)\n  return coll\n\n\n\ndef load_postproc_models(data):\n  models=[]\n  for i in range(5):\n    models.append(CorrInterp(data,n_hidden,n_hidden,REBINNINGS,device,ref=\"fit\"))\n    models[-1].load_state_dict(torch.load(\"/kaggle/input/ariel2025-postproc-nets-v19post/postproc_v19post_net_fold\"+str(i)+\".pth\",weights_only=True,map_location=device))\n    models[-1].eval()\n\n  fallbacks=[]\n  for i in range(5):\n    fallbacks.append(CorrInterp(data,n_hidden,n_hidden,REBINNINGS,device,ref=\"baseline\"))\n    fallbacks[-1].load_state_dict(torch.load(\"/kaggle/input/ariel2025-postproc-nets-v19post/postproc_fallback_v19post_net_fold\"+str(i)+\".pth\",weights_only=True,map_location=device))\n    fallbacks[-1].eval()\n    \n  return models,fallbacks\n\ndef combine_measurements(pred_list):\n  #expect pred list to be a list of (almost certainly two) tensors of shape (batch_size,2,n_wavelengths).\n  #Elements with that second index ==0 are predicted mean, and elements with second index ==1 are predicted uncertainty \n  means=[p[:,0,:] for p in pred_list]\n  sigmas=[p[:,1,:] for p in pred_list]\n  weights=[torch.where(sig**2>0,1/sig**2,0) for sig in sigmas]\n  sumweight=sum(weights)\n  weights=[w/sumweight for w in weights]\n  mean=sum([m*w for m,w in zip(means,weights)])\n  sigma=torch.where(sumweight**0.5>0,1./sumweight**0.5,1)\n\n  if difference_is_uncert:\n    mean_max=means[0]\n    mean_min=means[0]\n    for m in means[1:]:\n      mean_max=torch.maximum(mean_max,m)\n      mean_min=torch.minimum(mean_min,m)\n    delta=torch.abs(mean_max-mean_min)/2\n    sigma=(sigma**2+delta**2)**0.5\n  new_pred=torch.cat([torch.unsqueeze(mean,dim=1),torch.unsqueeze(sigma,dim=1)],dim=1)\n  return new_pred\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-09T03:06:17.09186Z","iopub.execute_input":"2025-10-09T03:06:17.092083Z","iopub.status.idle":"2025-10-09T03:06:17.126719Z","shell.execute_reply.started":"2025-10-09T03:06:17.092057Z","shell.execute_reply":"2025-10-09T03:06:17.125916Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# other utilities","metadata":{}},{"cell_type":"code","source":"def simple_outlier_exclusion(ch0_signal):\n  #outlier exclusion for fgs was done in load_planet(); only need to do airs here\n  ch0_signal=ch0_signal.detach().cpu().numpy()\n  ch0_slide=np.lib.stride_tricks.sliding_window_view(ch0_signal,(1,1,50))\n  ch0_slide_mean=np.squeeze(np.squeeze(np.mean(ch0_slide,axis=-1),axis=-1),axis=-1)\n  ch0_slide_std=np.squeeze(np.squeeze(np.std(ch0_slide,axis=-1),axis=-1),axis=-1)\n  ch0_slide_mean=np.pad(ch0_slide_mean,((0,0),(0,0),(25,24)),mode='edge')\n  ch0_slide_std=np.pad(ch0_slide_std,((0,0),(0,0),(25,24)),mode='edge')\n  dev=(ch0_signal-ch0_slide_mean)/ch0_slide_std\n  ch0_signal=torch.tensor(np.where(np.abs(dev)<3,ch0_signal,ch0_slide_mean)).to(device)\n\n  return ch0_signal\n\ndef initial_transit_bounds(fgs_signal,ch0_signal):\n  fgs_gaps_ok=True\n  try:\n    fgs_tstart_mean,fgs_tstart_sig,fgs_tend_mean,fgs_tend_sig,scalogram=find_transit_region(fgs_signal)\n    #print(\"fgs transit region: \"+str(fgs_tstart_mean)+\"+/-\"+str(fgs_tstart_sig)+\" to \"+str(fgs_tend_mean)+\"+/-\"+str(fgs_tend_sig),flush=True)\n  except Exception as e:\n    print(\"fgs call to find_transit_region raises exception: \"+str(e),flush=True)\n    fgs_gaps_ok=False\n    fgs_tstart_mean,fgs_tstart_sig,fgs_tend_mean,fgs_tend_sig=None,None,None,None\n\n  if fgs_tstart_sig is not None and fgs_tend_sig is not None and fgs_tstart_sig is not None and fgs_tend_sig is not None:\n    if fgs_tstart_sig<15 or fgs_tend_sig<15 or fgs_tstart_sig>1000 or fgs_tend_sig>1000:\n      fgs_gaps_ok=False\n\n  ch0_gaps_ok=False #take this out to save time\n  #ch0_gaps_ok=True\n  #try:\n  #  ch0_tstart_mean,ch0_tstart_sig,ch0_tend_mean,ch0_tend_sig,scalogram=find_transit_region(ch0_signal)\n  #except Exception as e:\n  #  print(\"ch0 call to find_transit_region raises exception: \"+str(e),flush=True)\n  #  ch0_gaps_ok=False\n  #  ch0_tstart_mean,ch0_tstart_sig,ch0_tend_mean,ch0_tend_sig=None,None,None,None#\n\n  #if ch0_tstart_mean is not None and ch0_tstart_sig is not None and ch0_tend_mean is not None and ch0_tend_sig is not None:\n  #  if ch0_tstart_sig<15 or ch0_tend_sig<15 or ch0_tstart_sig>1000 or ch0_tend_sig>1000:\n  #    ch0_gaps_ok=False\n\n  #if not ch0_gaps_ok and not fgs_gaps_ok:\n  #  tstart_mean=1800\n  #  tend_mean=3600\n  #  tstart_sig=50\n  #  tend_sig=50\n  #elif not ch0_gaps_ok:\n  #  if fgs_gaps_ok:\n  #    tstart_mean,tstart_sig,tend_mean,tend_sig=fgs_tstart_mean,fgs_tstart_sig,fgs_tend_mean,fgs_tend_sig\n  #elif (ch0_gaps_ok and fgs_gaps_ok) or (not ch0_gaps_ok and not fgs_gaps_ok):\n  #  tstart_mean=(ch0_tstart_mean+fgs_tstart_mean)/2\n  #  tstart_sig=(ch0_tstart_sig+fgs_tstart_sig)/2\n  #  tend_mean=(ch0_tend_mean+fgs_tend_mean)/2\n  #  tend_sig=(ch0_tend_sig+fgs_tend_sig)/2\n  #else:\n  #  tstart_mean,tstart_sig,tend_mean,tend_sig=ch0_tstart_mean,ch0_tstart_sig,ch0_tend_mean,ch0_tend_sig\n\n  if fgs_gaps_ok:\n    tstart_mean,tstart_sig,tend_mean,tend_sig=fgs_tstart_mean,fgs_tstart_sig,fgs_tend_mean,fgs_tend_sig\n  else:\n    tstart_mean=1800\n    tend_mean=3600\n    tstart_sig=50\n    tend_sig=50\n    \n  gap1=(int((tstart_mean-tstart_sig*turn_on_width)/fit_bin_width)*fit_bin_width,\n    fit_bin_width*(1+int((tstart_mean+tstart_sig*turn_on_width)/fit_bin_width)))\n  gap2=(int((tend_mean-tend_sig*turn_on_width)/fit_bin_width)*fit_bin_width,\n    fit_bin_width*(1+int((tend_mean+tend_sig*turn_on_width)/fit_bin_width)))\n\n  if gap2[0]<=gap1[1]+fit_bin_width:\n    midpoint=int((gap1[1]+gap2[0])/2)\n    gap1=(gap1[0],midpoint-fit_bin_width)\n    gap2=(midpoint+fit_bin_width,gap2[1])\n  if gap1[0]>=gap1[1]:\n    gap1=(max(0,gap1[1]-fit_bin_width),gap1[1])\n  if gap2[0]>=gap2[1]:\n    gap2=(gap2[0],min(ch0_signal.shape[-1],gap2[0]+fit_bin_width))\n\n      \n  return tstart_mean,tstart_sig,tend_mean,tend_sig, gap1, gap2\n\ndef smooth_errors(yerr):\n  pad_size=tuple([(0,0) for idim in range(len(yerr.shape)-1)]+[(SMOOTH_SIZE,SMOOTH_SIZE)])\n  yerr=np.lib.stride_tricks.sliding_window_view(np.pad(yerr.cpu().numpy(),pad_size,'edge'),2*SMOOTH_SIZE+1,axis=-1)\n  yerr=np.mean(yerr,axis=-1)*error_smoothing_scalefac\n  yerr=torch.tensor(yerr).to(device)\n  return yerr\n\ndef gaussian_error_propagation_baseline(y,yerr,gap1,gap2):\n  intrans=y[:,:,int(gap1[1]/fit_bin_width):max(int(gap1[1]/fit_bin_width)+1,int(gap2[0]/fit_bin_width))]   #expect shape (1,282,n_intrans)\n  outrans=torch.cat([y[:,:,0:max(1,int(gap1[0]/fit_bin_width))],y[:,:,min(y.shape[-1]-1,int(gap2[1]/fit_bin_width)):]],dim=-1)  #shape (1,282,n_outrans)\n  intrans_err=yerr[:,:,int(gap1[1]/fit_bin_width):max(int(gap1[1]/fit_bin_width)+1,int(gap2[0]/fit_bin_width))]   #expect shape (1,282,n_intrans)\n  outrans_err=torch.cat([yerr[:,:,0:max(1,int(gap1[0]/fit_bin_width))],yerr[:,:,min(yerr.shape[-1]-1,int(gap2[1]/fit_bin_width)):]],dim=-1)  #shape (1,282,n_outrans)\n\n  mean_intrans=torch.mean(intrans,dim=-1)\n  mean_outrans=torch.mean(outrans,dim=-1)\n  stderr_intrans=torch.std(intrans,dim=-1)/intrans.shape[-1]**0.5\n  stderr_outrans=torch.std(outrans,dim=-1)/outrans.shape[-1]**0.5\n  prop_intrans=torch.sqrt(torch.sum(intrans_err**2,dim=-1))\n  prop_outrans=torch.sqrt(torch.sum(outrans_err**2,dim=-1))\n  #dFF is short for \"delta F over F\"\n  dFF_baseline=(mean_outrans-mean_intrans)/mean_outrans\n  #propagation of error:  \n  #  d/d(mean_intrans) dFF = -1/mean_outrans\n  #  d/d(mean_outrans) dFF = mean_intrans/mean_outrans**2\n  #  sig(dFF)**2 = ((-1/mean_outrans)**2) d(mean_intrans)**2 + ((mean_intrans/mean_outrans**2)**2) d(mean_outrans)**2 \n  #dFF_baseline_err= ( ((-1/mean_outrans)**2) * (stderr_intrans**2 + prop_intrans**2) + ((mean_intrans/mean_outrans**2)**2) * (stderr_outrans**2 + prop_outrans**2))**0.5\n  dFF_baseline_err= ( ((-1/mean_outrans)**2) * (stderr_intrans**2 + prop_intrans**2) )**0.5\n\n  dFF_baseline=torch.flip(dFF_baseline,dims=(-1,))\n  dFF_baseline_err=torch.flip(dFF_baseline_err,dims=(-1,))\n\n  return dFF_baseline,dFF_baseline_err    \n\n\ndef compensate_limb_darkening_linear(fitted_depth,ldc_fit_param,b):\n  #just a quick-and-dirty ballpark estimate of the impact of limb darkening based on the fact that\n  #the difference in observed intensity at r=1 and at conjunction can be estimated from ldc_fit_param.\n  #It gets unstable when you have an impact parameter close to 1 and somehow also fit a large ldc parameter.\n  #Signal parameterization is a trapezoid; at conjunction (r=b) it is 1, while at ingress/egress it is 1-ldc (if we ignore tau),\n  #so the difference in intensity is just the value of the ldc parameter.\n  ldc=MAX_LDC_COEF*torch.sigmoid(torch.clip(torch.tensor(ldc_fit_param),-5,5)).cpu().numpy()\n  intensity_diff=ldc\n  #eq 8 in https://arxiv.org/pdf/1901.01730 : I = 1 - u*(1-(1-r**2)**0.5) --> 1-u at r=1\n  #so in terms of the limb darkening model, intensity_diff(@r=b - @r=1) =  [-u*(1-(1-b**2)**0.5)] +u = u*(1-b**2)**0.5, so that \n  u=max(0,min(1,intensity_diff/(1-b**2)**0.5))\n  #...where I have imposed by hand the limits on u so that we never have Ip going negative\n  Ia=1-u/3  #eq. 11\n  Ip=1-u*(1-(1-b**2)**0.5)  #eq.8\n  #if Ia/Ip<0:\n  #  print(\"compensate_limb_darkening_linear changes sign: Ia=\"+str(Ia)+\", Ip=\"+str(Ip)+\", u=\"+str(u)+\", b=\"+str(b)+\", ldc_fit_param=\"+str(ldc_fit_param),flush=True)\n  return fitted_depth*Ia/Ip\n\ndef load_bounds(fname):\n  out=dict()\n  with open(fname,\"r\") as f:\n    for line in f:\n      key,val=re.split(\",\",line[:-1])\n      if len(key)>0:\n        out[key]=float(val)\n  return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-09T03:06:17.127531Z","iopub.execute_input":"2025-10-09T03:06:17.12778Z","iopub.status.idle":"2025-10-09T03:06:17.150943Z","shell.execute_reply.started":"2025-10-09T03:06:17.127754Z","shell.execute_reply":"2025-10-09T03:06:17.150249Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# main program","metadata":{}},{"cell_type":"code","source":"try:\n  start_time=time_lib.time()\n  run_time_list=[]\n  load_time_list=[]\n  cpu_fit_time_list=[]\n  gpu_fit_time_list=[]\n  postproc_time_list=[]\n  mask_list=[]\n  print(\"started run at \"+str(start_time),flush=True)\n  meta_dict=load_metadata()\n  integration_time=meta_dict[\"integration_time\"]\n  cumulative_time=meta_dict[\"cumulative_time\"]\n\n  #planets=[k for k in meta_dict.keys() if \"time\" not in k]  <--may not be in the correct order\n  star_params,star_params_norm=load_star_params()\n\n  nn_input_var_min=load_bounds(\"/kaggle/input/ariel2025-nn-input-bounds-v19-debugncg/nn_input_min_bounds_v19_debugNCG.csv\")\n  nn_input_var_max=load_bounds(\"/kaggle/input/ariel2025-nn-input-bounds-v19-debugncg/nn_input_max_bounds_v19_debugNCG.csv\")\n    \n  error_code=None\n\n  #will load models after we have an example of the input\n  #...it was a convenient way to set things up during model development, albeit a bit awkward here\n  postproc_models=None\n  postproc_fallbacks=None\n  dFF_preds=[]\n  dFF_sigmas=[]\n  planets=meta_dict[\"index\"]\n  final_means=[]\n  final_sigmas=[]\n  last_example_start_time=time_lib.time()\n  for iplanet,planet in enumerate(planets):\n        \n    #if iplanet%10==0:\n    print(\"working on planet \"+str(planet)+\"; this is number \"+str(iplanet)+\" of \"+str(len(planets)),flush=True)\n\n    ch0_offset,ch0_gain,fgs_offset,fgs_gain,star=meta_dict[planet]\n\n    ls=os.listdir(\"/kaggle/input/ariel-data-challenge-2025/\"+str(dataset)+\"/\"+str(planet)+\"/\")\n    ls=[re.sub(\"FGS1_calibration_\",\"\",f) for f in ls if \"FGS1_calibration_\" in f]\n    #print(\"planet \"+str(planet)+\", ls=\"+str(ls),flush=True)  #expect either ['0'] or ['0','1'] -- maybe longer lists if the dataset has any examples of more than two visits per system\n\n    batch=[]\n    predicted_spectra=[]\n    for ivisit,visit in enumerate(ls):\n     try:\n      example_start_time=time_lib.time()\n      if iplanet>0 or ivisit>0:\n        run_time_list.append(example_start_time-last_example_start_time)\n        print(\"last example ran for \"+str(run_time_list[-1])+\"; average so far=\"+str(sum(run_time_list)/len(run_time_list)),flush=True)\n        last_example_start_time=example_start_time\n\n      ch0_signal,fgs_signal,ch0_signal_lincorr_shift,ch0_dsamp,fgs_nmask,ch0_nmask,fgs_mean_proj1, fgs_std_proj1, fgs_mean_proj2, fgs_std_proj2, ch0_mean, ch0_std=load_planet(planet,ch0_offset,ch0_gain,fgs_offset,fgs_gain,integration_time,visit)\n      ch0_signal=simple_outlier_exclusion(ch0_signal)\n      \n      fgs_signal=torch.reshape(fgs_signal,(fgs_signal.shape[0],fgs_signal.shape[1],-1,12))\n      fgs_signal=torch.mean(fgs_signal,dim=-1,keepdim=False)\n      ch0_signal=ch0_signal[:,drop_shoulders:-drop_shoulders,:] \n\n      ch0_mean=ch0_mean[drop_shoulders:-drop_shoulders,:]\n      ch0_std=ch0_std[drop_shoulders:-drop_shoulders,:]\n      fgs_mean_proj1=fgs_mean_proj1/32\n      fgs_mean_proj2=fgs_mean_proj2/32\n      fgs_std_proj1=fgs_std_proj1/32\n      fgs_std_proj2=fgs_std_proj2/32\n\n      ch0_mean=ch0_mean/32\n      ch0_std=ch0_std/32\n \n      ch0_spatial_center_mean=torch.mean(ch0_mean,dim=-1)\n      ch0_spatial_center_std=torch.std(ch0_mean,dim=-1)\n      ch0_spatial_width_mean=torch.mean(ch0_std,dim=-1)\n      ch0_spatial_width_std=torch.std(ch0_std,dim=-1)\n      ch0_spatial_center_range=torch.max(ch0_mean,dim=-1).values-torch.min(ch0_mean,dim=-1).values\n      ch0_spatial_width_range=torch.max(ch0_std,dim=-1).values-torch.min(ch0_std,dim=-1).values\n\n      tstart_mean,tstart_sig,tend_mean,tend_sig,gap1,gap2=initial_transit_bounds(fgs_signal,ch0_signal)\n      print(\"intial transit region: \"+str(tstart_mean)+\"+/-\"+str(tstart_sig)+\" to \"+str(tend_mean)+\"+/-\"+str(tend_sig),flush=True)\n\n      fit_outputs=[dict() for i in range(283)]\n      \n      fit_outputs[0][\"nmask_scaled\"]=fgs_nmask/(67500*32*32)\n      fit_outputs[0][\"spatial_center_mean\"]=(fgs_mean_proj1.mean().item()+fgs_mean_proj2.mean().item())/2\n      fit_outputs[0][\"spatial_center_std\"]=(fgs_mean_proj1.std().item()+fgs_mean_proj2.std().item())/2\n      fit_outputs[0][\"spatial_width_mean\"]=(fgs_std_proj1.mean().item()+fgs_std_proj2.mean().item())/2\n      fit_outputs[0][\"spatial_width_std\"]=(fgs_std_proj1.std().item()+fgs_std_proj2.std().item())/2\n      r1=torch.max(fgs_mean_proj1).item()-torch.min(fgs_mean_proj1).item()\n      r2=torch.max(fgs_mean_proj2).item()-torch.min(fgs_mean_proj2).item()\n      fit_outputs[0][\"spatial_center_range\"]=(r1+r2)/2\n\n      r1=torch.max(fgs_std_proj1).item()-torch.min(fgs_std_proj1).item()\n      r2=torch.max(fgs_std_proj2).item()-torch.min(fgs_std_proj2).item()\n      fit_outputs[0][\"spatial_width_range\"]=(r1+r2)/2\n\n      for j in range(282):\n        fit_outputs[j+1][\"nmask_scaled\"]=ch0_nmask[j]/(5625.*32.)\n        fit_outputs[j+1][\"spatial_center_mean\"]=ch0_spatial_center_mean[j]\n        fit_outputs[j+1][\"spatial_center_std\"]=ch0_spatial_center_std[j]\n        fit_outputs[j+1][\"spatial_width_mean\"]=ch0_spatial_width_mean[j]\n        fit_outputs[j+1][\"spatial_width_std\"]=ch0_spatial_width_std[j]\n        fit_outputs[j+1][\"spatial_center_range\"]=ch0_spatial_center_range[j]\n        fit_outputs[j+1][\"spatial_width_range\"]=ch0_spatial_width_range[j]\n\n      #validation printouts\n      #print(\"planet \"+str(planet)+\" fgs nmask_scaled=\"+str(fit_outputs[0][\"nmask_scaled\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" fgs spatial_center_mean=\"+str(fit_outputs[0][\"spatial_center_mean\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" fgs spatial_center_std=\"+str(fit_outputs[0][\"spatial_center_std\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" fgs spatial_center_range=\"+str(fit_outputs[0][\"spatial_center_range\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" fgs spatial_width_mean=\"+str(fit_outputs[0][\"spatial_width_mean\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" fgs spatial_width_std=\"+str(fit_outputs[0][\"spatial_width_std\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" fgs spatial_width_range=\"+str(fit_outputs[0][\"spatial_width_range\"]),flush=True)\n\n\n      #print(\"planet \"+str(planet)+\" ch0 150 nmask_scaled=\"+str(fit_outputs[151][\"nmask_scaled\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" ch0 150 spatial_center_mean=\"+str(fit_outputs[151][\"spatial_center_mean\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" ch0 150 spatial_center_std=\"+str(fit_outputs[151][\"spatial_center_std\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" ch0 150 spatial_center_range=\"+str(fit_outputs[151][\"spatial_center_range\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" ch0 150 spatial_width_mean=\"+str(fit_outputs[151][\"spatial_width_mean\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" ch0 150 spatial_width_std=\"+str(fit_outputs[151][\"spatial_width_std\"]),flush=True)\n      #print(\"planet \"+str(planet)+\" ch0 150 spatial_width_range=\"+str(fit_outputs[151][\"spatial_width_range\"]),flush=True)\n\n\n      fgs_mean_normalization=torch.mean(fgs_signal,dim=-1,keepdim=True)\n      fit_outputs[0][\"norm\"]=fgs_mean_normalization[0,0,:]\n\n      #print(\"planet \"+str(planet)+\" fgs norm=\"+str(fit_outputs[0][\"norm\"]),flush=True)\n        \n      fgs_norm=fgs_signal/fgs_mean_normalization\n      fgs_norm=torch.squeeze(fgs_norm,dim=1)\n\n      ch0_norm=torch.mean(ch0_signal,dim=-1,keepdim=True)\n      for j in range(ch0_signal.shape[1]):\n        fit_outputs[j+1][\"norm\"]=ch0_norm[0,j,0].item()\n\n      #print(\"planet \"+str(planet)+\" ch0 150 norm=\"+str(fit_outputs[151][\"norm\"]),flush=True)\n\n      for j in range(len(fit_outputs)):\n        fit_outputs[j][\"tstart_mean\"]=tstart_mean\n        fit_outputs[j][\"tstart_sig\"]=tstart_sig\n        fit_outputs[j][\"tend_mean\"]=tend_mean\n        fit_outputs[j][\"tend_sig\"]=tend_sig\n\n      ch0_signal_x=np.arange(0,1,1./ch0_signal.shape[-1])\n\n      #compute y and yerr values bins over time\n      ch0_signal_bins=torch.reshape(ch0_signal,tuple(list(ch0_signal.shape)[0:-1]+[-1,fit_bin_width]))\n      ch0_signal_y=torch.mean(ch0_signal_bins,dim=-1)\n      ch0_signal_yerr=torch.std(ch0_signal_bins,dim=-1)/fit_bin_width**0.5\n\n      #print(\"3\",flush=True)\n      fgs_bins=torch.reshape(fgs_signal,tuple(list(fgs_signal.shape)[0:-1]+[-1,fit_bin_width]))\n      fgs_y=torch.mean(fgs_bins,dim=-1)\n      fgs_yerr=torch.std(fgs_bins,dim=-1)/fit_bin_width**0.5\n      \n      if do_error_smoothing:\n        #smooth out the errors so that fluctuations don't give us very different errors in neighboring bins\n        ch0_signal_yerr=smooth_errors(ch0_signal_yerr)\n        fgs_yerr=smooth_errors(fgs_yerr)\n        \n      #print(\"ch0_signal_x.shape=\"+str(ch0_signal_x.shape),flush=True)\n      ch0_signal_x_binned=np.reshape(ch0_signal_x,(-1,fit_bin_width))\n      ch0_signal_x_binned=np.mean(ch0_signal_x_binned,axis=-1)\n\n      #print(\"4\",flush=True)\n      #a baseline check:  gaussian error propagation\n      dFF_baseline,dFF_baseline_err=gaussian_error_propagation_baseline(ch0_signal_y,ch0_signal_yerr,gap1,gap2)\n      dFF_fgs_baseline,dFF_fgs_baseline_err=gaussian_error_propagation_baseline(fgs_y,fgs_yerr,gap1,gap2)\n\n      #print(\"dFF_baseline.shape=\"+str(dFF_baseline.shape),flush=True)  #(1,len(freq_bins))\n      #print(\"dFF_fgs_baseline.shape=\"+str(dFF_fgs_baseline.shape),flush=True)  #(1,1)\n\n      dFF_preds.append(torch.cat([dFF_fgs_baseline,dFF_baseline],dim=-1).detach().cpu().numpy())\n      dFF_sigmas.append(torch.cat([dFF_fgs_baseline_err,dFF_baseline_err],dim=-1).detach().cpu().numpy())\n      \n      fit_outputs[0][\"baseline_depth\"]=dFF_fgs_baseline[0,0]\n      fit_outputs[0][\"baseline_depth_err\"]=dFF_fgs_baseline_err[0,0]\n      for j in range(282):\n        fit_outputs[j+1][\"baseline_depth\"]=dFF_baseline[0,j]\n        fit_outputs[j+1][\"baseline_depth_err\"]=dFF_baseline_err[0,j]\n\n      #print(\"planet \"+str(planet)+\" fgs baseline_depth=\"+str(fit_outputs[0][\"baseline_depth\"].item()),flush=True)\n      #print(\"planet \"+str(planet)+\" ch0 150 baseline_depth=\"+str(fit_outputs[151][\"baseline_depth\"].item()),flush=True)\n\n      #print(\"planet \"+str(planet)+\" fgs baseline_depth_err=\"+str(fit_outputs[0][\"baseline_depth_err\"].item()),flush=True)\n      #print(\"planet \"+str(planet)+\" ch0 150 baseline_depth_err=\"+str(fit_outputs[151][\"baseline_depth_err\"].item()),flush=True)\n\n            \n      #set up some initial guess values for fit parameters\n      Tcenter_guess=abs((sum(gap1)+sum(gap2))/4.)/fit_bin_width\n      T_full_guess=abs(gap2[0]-gap1[1])/fit_bin_width\n      T_tot_guess=abs(gap2[1]-gap1[0])/fit_bin_width\n      T_guess=(T_full_guess+T_tot_guess)/2\n\n      #print(\"gap1=\"+str(gap1),flush=True)\n      #print(\"gap2=\"+str(gap2),flush=True)\n      #print(\"T_full_guess=\"+str(T_full_guess)+\", T_tot_guess=\"+str(T_tot_guess)+\", T_guess=\"+str(T_guess),flush=True)\n\n      tau_guess=(T_tot_guess-T_full_guess)/2\n      #print(\"tau_guess (step 1): \"+str(tau_guess),flush=True)\n\n      tau_guess=min(T_guess/2-1,tau_guess)/(T_guess/2)\n      #print(\"tau_guess (step 2): \"+str(tau_guess),flush=True)\n\n      tau_guess=-np.log((1./tau_guess)-1)\n      #print(\"tau_guess (step 3): \"+str(tau_guess),flush=True)\n\n      intransit_start=int(gap1[1]/fit_bin_width)\n      intransit_end=int(gap2[0]/fit_bin_width)\n      if intransit_end<=intransit_start:\n        midpoint=int((intransit_start+intransit_end)/2)\n        intransit_start=midpoint\n        intransit_end=midpoint+1\n\n      #print(\"planet in star_params? \"+str(planet in star_params),flush=True)\n      Rs,Ms,Ts,Mp,e,P,sma,i=star_params[planet]\n      #i=i*np.pi/180  #i is given in degrees; most trig functions want radians  <--already applied in load_star_params\n      w=0  #...but for circular orbits, it shouldn't matter, since terms involving w always get multiplied by e\n      fac16=(1-e*e)**0.5/(1+e*np.sin(w))  #Winn eq 16\n      b=(sma*np.cos(i))*((1-e*e)/(1+e*np.sin(w)))  #I think sma is already the ratio semi-major-axis/Rs\n      #print(\"predicted impact parameter for this transit: \"+str(b),flush=True)\n      if b>0.9999:  #well, b>1, but with a bit of wiggle room\n        #print(\"unphysical impact parameter: \"+str(b),flush=True)\n        #print(\"sma=\"+str(sma)+\", inclination=\"+str(i),flush=True)\n        #print(\"e=\"+str(e)+\", w=\"+str(w),flush=True)\n        b=0.9999\n\n      #k=Rp/Rs  #--> you do not haz; cannot compute this outside of a fit where Rp is a free parameter\n      k=0  #-->compute these preds in the limit Rp->0 ==> k->0\n      T_tot_pred=fac16*(P/np.pi)*np.arcsin((1/sma)*(((1+k)**2-b**2)**0.5/(np.sin(i))))  #sma is (maybe) already semi-major-axis/Rs\n      T_full_pred=fac16*(P/np.pi)*np.arcsin((1/sma)*(((1-k)**2-b**2)**0.5/(np.sin(i))))\n      #print(\"Detected from data and expressed in hours, T_full=\"+str(T_full_guess)+\" and T_tot=\"+str(T_tot_guess),flush=True)\n\n      #the above are expected to be in units of hours, so convert to bins\n      T_tot_pred=T_tot_pred*(1/7.5)*ch0_signal_x_binned.shape[-1]\n      T_full_pred=T_full_pred*(1/7.5)*ch0_signal_x_binned.shape[-1]\n\n      #print(\"Detected from data and expressed in bins, T_full=\"+str(T_full_guess)+\" and T_tot=\"+str(T_tot_guess)+\"; mean=\"+str((T_full_guess+T_tot_guess)/2),flush=True)\n      #print(\"Predicted by calculations: T_full=\"+str(T_full_pred)+\" and T_tot=\"+str(T_tot_pred),flush=True)\n\n      T_tot_pred*=fit_bin_width\n      T_full_pred*=fit_bin_width\n      #print(\"Predicted by calculations and multiplying by rebin factor of \"+str(fit_bin_width)+\": T_full=\"+str(T_full_pred)+\" and T_tot=\"+str(T_tot_pred),flush=True)\n\n      fgsn=torch.mean(fgs_y,dim=-1,keepdim=True)\n      fgs_y=fgs_y/fgsn\n      fgs_yerr=fgs_yerr/fgsn\n\n      #print(\"intransit_start=\"+str(intransit_start)+\", intransit_end=\"+str(intransit_end),flush=True)\n      intransit_max=torch.max(fgs_y[...,intransit_start:intransit_end]).detach().cpu().numpy()\n      intransit_min=torch.min(fgs_y[...,intransit_start:intransit_end]).detach().cpu().numpy()\n      intransit_range=intransit_max-intransit_min\n      intransit_range=intransit_max-intransit_min\n      ldc_guess=(intransit_range/2)**0.5\n      if ldc_guess<0.001 or ldc_guess>MAX_LDC_COEF*0.999:\n        #print(\"warning: ldc_guess out of range: \"+str(ldc_guess)+\"; MAX_LDC_COEF=\"+str(MAX_LDC_COEF),flush=True)\n        ldc_guess=max(0.001,min(0.999*MAX_LDC_COEF,ldc_guess))\n      #also convert to logits, since that is what we will fit\n      ldc_guess=np.log((MAX_LDC_COEF/ldc_guess)-1)\n\n      #print(\"guess params: \"+str([Tcenter_guess,T_guess,tau_guess,ldc_guess]),flush=True)\n\n      cpu_fit_start_time=time_lib.time()\n\n      #First, fit the FGS light curve, and then use parameters from that fit to initialize/guide the fit to the AIRS data\n      fit_funcs=[]\n      penalty_funcs=[]\n\n      #parameter list:\n      #  - signal norm: [0.]\n      #  - background shape: [0]*NPARAM-1 + [1.], where the last entry is the constant term in the polynomial\n      #  - signal center logit: [0]\n      #  - transit duration: [T_guess]\n      #  - ingress/egress duration logit: [-2]\n      #  - limb darkening coeffs: [0.,0.]  --> [0.] if we choose not to fit x**4 term\n      #init_data.append(np.array(tuple([0 for i in range(NPARAM)]+[1.,np.log((ch0_signal_x_binned.shape[-1]/Tcenter_guess)-1),T_guess,-2,ldc_guess]),dtype=np.float64))\n      #init_data=np.array(tuple([0 for i in range(NPARAM)]+[1.,0.,0.,tau_guess,ldc_guess]),dtype=np.float64)\n      init_depth=np.log(np.exp(100*dFF_fgs_baseline[0,0].item())-1)/100.\n      init_data=np.array(tuple([init_depth]+[0 for i in range(NPARAM-1)]+[1.,0.,0.,tau_guess,ldc_guess]),dtype=np.float64)\n\n      #print(\"for cpu fit, init_data=\"+str(init_data),flush=True)\n      fit_funcs.append(compute_polynomial_fit)\n      penalty_funcs.append(polynomial_penalty)\n\n      #print(\"fgs_y=\"+str(fgs_y),flush=True)\n        \n      data=((ch0_signal_x_binned,Tcenter_guess,T_guess),torch.squeeze(fgs_y[0,:]).detach().cpu().numpy(),torch.squeeze(fgs_yerr[0,:]).detach().cpu().numpy(),fit_funcs[0],penalty_funcs[0])\n      arg_tup=(light_curve_chisquare,init_data,data,{})\n      #print(\"calling cpu_fit\",flush=True)\n      fgs_central,hess_uncerts=cpu_fit(arg_tup,is_retry=False)\n      #print(\"back from cpu_fit\",flush=True)\n      #print(\"fgs_central.fun=\"+str(fgs_central.fun),flush=True)\n      print(\"fgs_central.x=\"+str(fgs_central.x),flush=True)\n      #print(\"fgs_y.device=\"+str(fgs_y.device),flush=True)\n      chisq_ndof=(fgs_central.fun-penalty_funcs[0](fgs_central.x))/(fgs_y.shape[-1]-(len(init_data)-1))\n      #print(\"chisq_ndof=\"+str(chisq_ndof),flush=True)\n        \n      #if len(hess_uncerts)==0:\n      #  print(\"no hess_uncerts from first fit; try with stricter constraints on some parameters\",flush=True)\n      #  #try a stricter penalty function that puts more weight on the initial guesses\n      #  data_check=((ch0_signal_x_binned,Tcenter_guess,T_guess),torch.squeeze(fgs_y[0,:]).detach().cpu().numpy(),torch.squeeze(fgs_yerr[0,:]).detach().cpu().numpy(),fit_funcs[0],polynomial_penalty)\n      #  arg_tup=(light_curve_chisquare,init_data,data_check,lowerbound,upperbound,0,{})\n      #  #print(\"calling cpu_fit with is_retry=True\",flush=True)\n      #  fgs_lower_bound,fgs_upper_bound,fgs_central,hess_uncerts=cpu_fit(arg_tup,is_retry=True)\n      #  print(\"after second attempt, fgs_central.fun=\"+str(fgs_central.fun),flush=True)\n\n      #columns in avg_hess_uncerts:\n      #wavelength_bin_id,raw_fit_depth_hess_err,lincorr_fit_depth_err,Tcenter_hess_err,T_hess_err,tau_hess_err,ldc_hess_err,fit_p0_hess_err,fit_p1_hess_err,fit_p2_hess_err,fit_p3_hess_err\n      #--> have to make sure ordering is right\n      #--> each row corresponds to a wavelength; this is for fgs fit which is wavelength 0\n      #fit parameters (in order):\n      #  depth, p3, p2, p1, p0, Tcenter, T, tau, ldc\n      fit_errnames=[\n        \"raw_fit_depth_hess_err\",\n        \"fit_p3_hess_err\",\n        \"fit_p2_hess_err\",\n        \"fit_p1_hess_err\",\n        \"fit_p0_hess_err\",\n        \"Tcenter_hess_err\",\n        \"T_hess_err\",\n        \"tau_hess_err\",\n        \"ldc_hess_err\"\n      ]\n        \n      if len(hess_uncerts)==0:\n        #hess_uncerts=avg_hess_uncerts.iloc[0].tolist()[1:]\n        hess_uncerts=[avg_hess_uncerts.iloc[0][name] for name in fit_errnames]\n        print(\"still no hess_uncerts, so we have to fall back to default values for uncertainties\",flush=True)\n\n      #the above handles the case where there are *no* uncertainties returned by the fit, but it may also be the case\n      #that a basically-failed error analysis returns a garbage value (10k, which is 1/sqrt(eps) where eps is a \n      #regularization factor added to the diagonal of the hessian).  Impute average values for these cases as well.\n      for j in range(len(hess_uncerts)):\n        for ivar,name in enumerate(fit_errnames):\n          limit=1 #all these errors are normally tiny except for the ones that are 10k\n          if name==\"ldc_hess_err\":  #except this one, which is not always quite so tiny\n            limit=100\n          if hess_uncerts[ivar]>limit:\n            print(\"hess_uncerts above limit for ivar=\"+str(ivar)+\", name=\"+str(name),flush=True)\n            hess_uncerts[ivar]=avg_hess_uncerts.iloc[0][name]\n\n      #print(\"all done with fgs fit\",flush=True)\n\n      chisq_ndof=(fgs_central.fun-penalty_funcs[0](fgs_central.x))/(fgs_y.shape[-1]-(len(init_data)-1))\n      print(\"chisq_ndof=\"+str(chisq_ndof),flush=True)\n\n      fgs_cv=torch.nn.functional.softplus(torch.tensor(fgs_central.x[0]),beta=100)\n      if len(hess_uncerts)>0:\n        fgs_err=hess_uncerts[0]\n      else:\n        fgs_err=0\n\n      fit_outputs[0][\"chisq_ndof\"]=chisq_ndof\n      fit_outputs[0][\"raw_fit_depth\"]=fgs_cv\n      if len(hess_uncerts)>0:\n        fit_outputs[0][\"raw_fit_depth_hess_err\"]=hess_uncerts[0]\n      else:\n        fit_outputs[0][\"raw_fit_depth_hess_err\"]=0\n\n      fgs_cv_lincorr=compensate_limb_darkening_linear(fgs_cv,fgs_central.x[-1],b)\n      fgs_lowerbound_lincorr=compensate_limb_darkening_linear(fgs_cv-fgs_err,fgs_central.x[-1],b)\n      fgs_upperbound_lincorr=compensate_limb_darkening_linear(fgs_cv+fgs_err,fgs_central.x[-1],b)\n      fgs_err_lincorr=(abs(fgs_cv_lincorr-fgs_lowerbound_lincorr)+abs(fgs_upperbound_lincorr-fgs_cv_lincorr))/2\n      print(\"corrected depth (fgs): \"+str(fgs_cv_lincorr)+\"+/-\"+str(fgs_err_lincorr),flush=True)\n\n      print(\"...and the fitted (fgs) transit depth is \"+str(fgs_cv)+\" +/-\"+str(fgs_err),flush=True)\n\n      fit_outputs[0][\"lincorr_fit_depth\"]=fgs_cv_lincorr.item()\n      fit_outputs[0][\"lincorr_fit_depth_err\"]=fgs_err_lincorr.item()\n      fit_outputs[0][\"Tcenter\"]=torch.clip(torch.tensor(Tcenter_guess+(fgs_y.shape[-1]/2)*fgs_central.x[-4]),0,fgs_y.shape[-1]).item()\n      fit_outputs[0][\"T\"]=torch.clip(T_guess+(fgs_y.shape[-1]/2)*torch.nn.functional.softplus(torch.tensor(fgs_central.x[-3]),beta=100),0,fgs_y.shape[-1])\n      fit_outputs[0][\"tau\"]=(fit_outputs[0][\"T\"]/2)*torch.sigmoid(torch.clip(torch.tensor(fgs_central.x[-2]),-5,5)).item()\n      fit_outputs[0][\"ldc\"]=MAX_LDC_COEF*torch.sigmoid(torch.clip(torch.tensor(fgs_central.x[-1]),-5,5)).item()\n      fit_outputs[0][\"fit_p0\"]=fgs_central.x[4]\n      fit_outputs[0][\"fit_p1\"]=fgs_central.x[3]\n      fit_outputs[0][\"fit_p2\"]=fgs_central.x[2]\n      fit_outputs[0][\"fit_p3\"]=fgs_central.x[1]\n      if len(hess_uncerts)>0:\n        fit_outputs[0][\"Tcenter_hess_err\"]=hess_uncerts[5]\n        fit_outputs[0][\"T_hess_err\"]=hess_uncerts[6]\n        fit_outputs[0][\"tau_hess_err\"]=hess_uncerts[7]\n        fit_outputs[0][\"ldc_hess_err\"]=hess_uncerts[8]\n        fit_outputs[0][\"fit_p3_hess_err\"]=hess_uncerts[1]\n        fit_outputs[0][\"fit_p2_hess_err\"]=hess_uncerts[2]\n        fit_outputs[0][\"fit_p1_hess_err\"]=hess_uncerts[3]\n        fit_outputs[0][\"fit_p0_hess_err\"]=hess_uncerts[4]\n      else:\n        fit_outputs[0][\"Tcenter_hess_err\"]=0\n        fit_outputs[0][\"T_hess_err\"]=0\n        fit_outputs[0][\"tau_hess_err\"]=0\n        fit_outputs[0][\"ldc_hess_err\"]=0\n        fit_outputs[0][\"fit_p3_hess_err\"]=0\n        fit_outputs[0][\"fit_p2_hess_err\"]=0\n        fit_outputs[0][\"fit_p1_hess_err\"]=0\n        fit_outputs[0][\"fit_p0_hess_err\"]=0\n\n      #print(\"1\",flush=True)\n      #pre_start=0\n      #pre_end=max(1,int(fit_outputs[0][\"Tcenter\"]-fit_outputs[0][\"T\"]/2-fit_outputs[0][\"tau\"]/2))\n      #ingr_start=max(pre_end+1,int(fit_outputs[0][\"Tcenter\"]-fit_outputs[0][\"T\"]/2+fit_outputs[0][\"tau\"]/2))\n      #ingr_end=ingr_start+max(1,int((fit_outputs[0][\"T\"]-fit_outputs[0][\"tau\"])/4))\n      #mid_start=int(fit_outputs[0][\"Tcenter\"]-((fit_outputs[0][\"T\"]-fit_outputs[0][\"tau\"])/4))\n      #mid_end=max(mid_start+1,int(fit_outputs[0][\"Tcenter\"]+((fit_outputs[0][\"T\"]-fit_outputs[0][\"tau\"])/4)))\n      #egr_start=mid_end\n      #egr_end=max(egr_start+1,int(fit_outputs[0][\"Tcenter\"]+fit_outputs[0][\"T\"]/2-fit_outputs[0][\"tau\"]/2))\n      #post_start=min(fgs_y.shape[-1]-1,int(fit_outputs[0][\"Tcenter\"]+fit_outputs[0][\"T\"]/2+fit_outputs[0][\"tau\"]/2))\n\n      pre_start=0\n      pre_end=max(1,int(fit_outputs[0][\"Tcenter\"]-fit_outputs[0][\"T\"]/2-fit_outputs[0][\"tau\"]/2))\n      ingr_start=max(0,int(fit_outputs[0][\"Tcenter\"]-fit_outputs[0][\"T\"]/2+fit_outputs[0][\"tau\"]/2))\n      ingr_end=ingr_start+max(1,int((fit_outputs[0][\"T\"]-fit_outputs[0][\"tau\"])/4))\n      mid_start=int(fit_outputs[0][\"Tcenter\"]-((fit_outputs[0][\"T\"]-fit_outputs[0][\"tau\"])/4))\n      mid_end=max(mid_start+1,int(fit_outputs[0][\"Tcenter\"]+((fit_outputs[0][\"T\"]-fit_outputs[0][\"tau\"])/4)))\n\n      egr_start=int(fit_outputs[0][\"Tcenter\"]+((fit_outputs[0][\"T\"]-fit_outputs[0][\"tau\"])/4))\n      egr_end=max(egr_start+1,int(fit_outputs[0][\"Tcenter\"]+fit_outputs[0][\"T\"]/2-fit_outputs[0][\"tau\"]/2))\n      post_start=min(fgs_y.shape[-1]-2,int(fit_outputs[0][\"Tcenter\"]+fit_outputs[0][\"T\"]/2+fit_outputs[0][\"tau\"]/2))\n\n\n      pre_end=min(pre_end,fgs_y.shape[-1]-1)\n      ingr_start=min(ingr_start,fgs_y.shape[-1]-2)\n      ingr_end=min(ingr_end,fgs_y.shape[-1]-1)\n      mid_start=min(max(mid_start,0),fgs_y.shape[-1]-2)\n      mid_end=min(max(mid_end,1),fgs_y.shape[-1]-1)\n      egr_start=min(max(egr_start,0),fgs_y.shape[-1]-2)\n      egr_end=min(egr_end,fgs_y.shape[-1]-1)\n\n        \n      #print(\"2\",flush=True)\n      #print(\"fgs_central.x=\"+str(fgs_central.x),flush=True)\n      #print(\"other compute_poly args: \"+str([Tcenter_guess,T_guess]))\n      fit_y=compute_polynomial_fit(fgs_central.x,(ch0_signal_x_binned,Tcenter_guess,T_guess),include_signal=True)\n      fit_y_nosig=compute_polynomial_fit(fgs_central.x,(ch0_signal_x_binned,Tcenter_guess,T_guess),include_signal=False)\n      fit_y_sigonly=fit_y_nosig-fit_y\n      #print(\"back from compute_polynomial_fit calls; fgs_y device=\"+str(fgs_y.device),flush=True)\n      residual=(fgs_y.to(cpu_fit_device)-fit_y)/fgs_yerr.to(cpu_fit_device)\n        \n      #print(\"3\",flush=True)\n      if pre_end>pre_start:\n        fit_outputs[0][\"res_pre\"]=residual[0,0,pre_start:pre_end].mean()\n        fit_outputs[0][\"bg_pre\"]=fit_y_nosig[pre_start:pre_end].mean()\n        fit_outputs[0][\"sig_pre\"]=fit_y_sigonly[pre_start:pre_end].mean()\n      else:\n        fit_outputs[0][\"res_pre\"]=0\n        fit_outputs[0][\"bg_pre\"]=0\n        fit_outputs[0][\"sig_pre\"]=0\n      if ingr_end>ingr_start:\n        fit_outputs[0][\"res_ingr\"]=residual[0,0,ingr_start:ingr_end].mean()\n        fit_outputs[0][\"bg_ingr\"]=fit_y_nosig[ingr_start:ingr_end].mean()\n        fit_outputs[0][\"sig_ingr\"]=fit_y_sigonly[ingr_start:ingr_end].mean()\n      else:\n        fit_outputs[0][\"res_ingr\"]=0\n        fit_outputs[0][\"bg_ingr\"]=0\n        fit_outputs[0][\"sig_ingr\"]=0\n      if mid_end>mid_start:\n        fit_outputs[0][\"res_mid\"]=residual[0,0,mid_start:mid_end].mean()\n        fit_outputs[0][\"bg_mid\"]=fit_y_nosig[mid_start:mid_end].mean()\n        fit_outputs[0][\"sig_mid\"]=fit_y_sigonly[mid_start:mid_end].mean()\n      else:\n        fit_outputs[0][\"res_mid\"]=0\n        fit_outputs[0][\"bg_mid\"]=0\n        fit_outputs[0][\"sig_mid\"]=0\n      if egr_end>egr_start:\n        fit_outputs[0][\"res_egr\"]=residual[0,0,egr_start:egr_end].mean()\n        fit_outputs[0][\"bg_egr\"]=fit_y_nosig[egr_start:egr_end].mean()\n        fit_outputs[0][\"sig_egr\"]=fit_y_sigonly[egr_start:egr_end].mean()\n      else:\n        fit_outputs[0][\"res_egr\"]=0\n        fit_outputs[0][\"bg_egr\"]=0\n        fit_outputs[0][\"sig_egr\"]=0\n      if residual.shape[-1]>post_start:\n        fit_outputs[0][\"res_post\"]=residual[0,0,post_start:].mean()\n        fit_outputs[0][\"bg_post\"]=fit_y_nosig[post_start:].mean()\n        fit_outputs[0][\"sig_post\"]=fit_y_sigonly[post_start:].mean()\n      else:\n        fit_outputs[0][\"res_post\"]=0\n        fit_outputs[0][\"bg_post\"]=0\n        fit_outputs[0][\"sig_post\"]=0\n\n      #print(\"4\",flush=True)\n      if normalize_batch_fits:\n        norm=torch.mean(ch0_signal_y,dim=-1,keepdims=True)\n        ch0_signal_y=ch0_signal_y/norm\n        ch0_signal_yerr=ch0_signal_yerr/norm\n\n      init_norm=torch.mean(ch0_signal_y[0,:,:],dim=-1,keepdims=False).to(device)\n      #print(\"fgs_central.x=\"+str(fgs_central.x),flush=True)\n\n      gpu_fit_start_time=time_lib.time()\n\n      init_depth=np.log(np.exp(100*fgs_cv.item())-1)/100.\n      init_data=[init_depth*torch.ones_like(init_norm).to(device), fgs_central.x[-2]*torch.ones_like(init_norm).to(device),fgs_central.x[-1]*torch.ones_like(init_norm).to(device)]\n      init_data=init_data+[init_norm]+[torch.zeros_like(init_norm).to(device) for i in range(NPARAM-1)]\n\n      #print(\"init_data for gpu fit: \"+str(init_data),flush=True)\n        \n      data=(torch.tensor(ch0_signal_x_binned,dtype=torch.float32,device=device),\n            ch0_signal_y[0,:,:].to(device),\n            ch0_signal_yerr[0,:,:].to(device),\n           (Tcenter_guess,T_guess,tau_guess,MAX_LDC_COEF,torch.tensor(fgs_central.x[-4:],device=device)))\n\n      #print(\"data for gpu fit: \"+str(data),flush=True)\n        \n      #print(\"init_data devices: \"+str([x.device for x in init_data]),flush=True)\n        \n      ok_poly,central_poly,ch0_hess_errs=gpu_fit(LightCurveChisquare,\n                                                 init_data,\n                                                 data,\n                                                 device,\n                                                 central=None,\n                                                 niter=2000, \n                                                 niter_errfit=50,  \n                                                 lr=4e-3, #9e-3,  #3e-3, \n                                                 lr_decay=1, \n                                                 eps=0.0003,noisy=False,force_refit=False)\n\n      #print(\"ok_poly=\"+str(ok_poly),flush=True)\n      #print(\"len(ch0_hess_errs)=\"+str(len(ch0_hess_errs)),flush=True)\n      if len(ch0_hess_errs)>0:  \n        #print(\"ch0_hess_errs 5=\"+str(ch0_hess_errs[:,5].tolist()),flush=True)\n        #print(\"...and from fgs, hess_uncerts=\"+str(hess_uncerts),flush=True)\n        #columns in avg_hess_uncerts:\n        #wavelength_bin_id,raw_fit_depth_hess_err,lincorr_fit_depth_err,Tcenter_hess_err,T_hess_err,tau_hess_err,ldc_hess_err,fit_p0_hess_err,fit_p1_hess_err,fit_p2_hess_err,fit_p3_hess_err\n        #fit errors (in order):\n        #  raw_fit_depth_hess_err,tau_hess_err, ldc_hess_err,fit_p0_hess_err,fit_p1_hess_err,fit_p2_hess_err,fit_p3_hess_err\n        avg_np=torch.tensor(avg_hess_uncerts.values)\n        #print(\"avg_np.shape=\"+str(avg_np.shape),flush=True)  #expect (283,11)\n        ch0_hess_errs[:,0]=torch.where(ch0_hess_errs[:,0]>1,avg_np[1:,1],ch0_hess_errs[:,0])\n        ch0_hess_errs[:,1]=torch.where(ch0_hess_errs[:,1]>1,avg_np[1:,5],ch0_hess_errs[:,1])\n        ch0_hess_errs[:,2]=torch.where(ch0_hess_errs[:,2]>100,avg_np[1:,6],ch0_hess_errs[:,2])\n        ch0_hess_errs[:,3]=torch.where(ch0_hess_errs[:,3]>1,avg_np[1:,7],ch0_hess_errs[:,3])\n\n        #these are a bit silly in v19_debugNCG because some fits do return uncertainties larger than 1.\n        #But due to an oversight, this is what was done during training, so replicate it here.\n        ch0_hess_errs[:,4]=torch.where(ch0_hess_errs[:,4]>1,avg_np[1:,8],ch0_hess_errs[:,4]) #*100)\n        ch0_hess_errs[:,5]=torch.where(ch0_hess_errs[:,5]>1,avg_np[1:,9],ch0_hess_errs[:,5]) #*100)\n        ch0_hess_errs[:,6]=torch.where(ch0_hess_errs[:,6]>1,avg_np[1:,10],ch0_hess_errs[:,6])  #*100)\n        #print(\"finished ch0_hess_errs imputations, now ch0_hess_errs=\"+str(ch0_hess_errs[:,5].tolist()),flush=True)\n\n      for j in range(282):\n        fit_outputs[j+1][\"chisq_ndof\"]=central_poly.fun[j]/(ch0_signal_y.shape[-1]-(len(init_data)-1))\n   \n      fit_results=[\n          ChisqFitResult(\n              success=ok_poly,\n              fun=central_poly.fun[i],\n              x=[x[i,...].item() for x in central_poly.x],\n              instance=central_poly.instance) for i in range(init_norm.shape[0])\n          ]\n  \n\n      for i in range(ch0_signal_y.shape[1]):\n        idx=i+1\n        central_poly=fit_results[i]\n        cv_raw=torch.nn.functional.softplus(torch.tensor(central_poly.x[0]),beta=100)\n    \n        fit_outputs[idx][\"raw_fit_depth\"]=cv_raw.item()\n        if len(ch0_hess_errs)>0:\n          fit_outputs[idx][\"raw_fit_depth_hess_err\"]=ch0_hess_errs[i,0]\n\n          cl_lower_bound_poly=cv_raw.item()-ch0_hess_errs[i,0]\n          cl_upper_bound_poly=cv_raw.item()+ch0_hess_errs[i,0]\n        else:\n          cl_upper_bound_poly=cl_lower_bound_poly=0\n          err_raw=0\n          fit_outputs[idx][\"raw_fit_depth_hess_err\"]=err_raw\n\n        err=ch0_hess_errs[i,0]\n\n        cv_lincorr=compensate_limb_darkening_linear(cv_raw,central_poly.x[2],b)\n        cv_lowerbound_lincorr=compensate_limb_darkening_linear(cv_raw-err,central_poly.x[2],b)\n        cv_upperbound_lincorr=compensate_limb_darkening_linear(cv_raw+err,central_poly.x[2],b)\n        err_lincorr=(abs(cv_lincorr-cv_lowerbound_lincorr)+abs(cv_upperbound_lincorr-cv_lincorr))/2\n    \n        if torch.is_tensor(err):\n          err=err.item()\n            \n        fit_outputs[idx][\"lincorr_fit_depth\"]=cv_lincorr.item()\n        fit_outputs[idx][\"lincorr_fit_depth_err\"]=err_lincorr\n        fit_outputs[idx][\"Tcenter\"]=torch.clip(Tcenter_guess+(fgs_y.shape[-1]/2)*torch.tensor(fgs_central.x[-4]),0,fgs_y.shape[-1]).item()\n        fit_outputs[idx][\"T\"]=torch.clip(T_guess+(fgs_y.shape[-1]/2)*torch.nn.functional.softplus(torch.tensor(fgs_central.x[-3]),beta=100),0,fgs_y.shape[-1])\n        fit_outputs[idx][\"tau\"]=(fit_outputs[idx][\"T\"]/2)*torch.sigmoid(torch.clip(torch.tensor(central_poly.x[1]),-5,5)).item()\n        fit_outputs[idx][\"ldc\"]=MAX_LDC_COEF*torch.sigmoid(torch.clip(torch.tensor(central_poly.x[2]),-5,5)).item()\n\n        fit_outputs[idx][\"fit_p0\"]=central_poly.x[3]\n        fit_outputs[idx][\"fit_p1\"]=central_poly.x[4] #*100\n        fit_outputs[idx][\"fit_p2\"]=central_poly.x[5] #*100\n        fit_outputs[idx][\"fit_p3\"]=central_poly.x[6] #*100\n\n        if len(ch0_hess_errs)>0:\n          fit_outputs[idx][\"tau_hess_err\"]=ch0_hess_errs[i,1] #in principle should transform this to match tau, but this is just a neural net input so let's not bother\n          fit_outputs[idx][\"ldc_hess_err\"]=ch0_hess_errs[i,2]\n          fit_outputs[idx][\"fit_p0_hess_err\"]=ch0_hess_errs[i,3]  #*100 has already happened above, when we did imputation of\n          fit_outputs[idx][\"fit_p1_hess_err\"]=ch0_hess_errs[i,4]  #average values to replace values of 10k from failed error analysis\n          fit_outputs[idx][\"fit_p2_hess_err\"]=ch0_hess_errs[i,5]\n          fit_outputs[idx][\"fit_p3_hess_err\"]=ch0_hess_errs[i,6]\n        else:\n          fit_outputs[idx][\"tau_hess_err\"]=0\n          fit_outputs[idx][\"ldc_hess_err\"]=0\n          fit_outputs[idx][\"fit_p0_hess_err\"]=0\n          fit_outputs[idx][\"fit_p1_hess_err\"]=0\n          fit_outputs[idx][\"fit_p2_hess_err\"]=0\n          fit_outputs[idx][\"fit_p3_hess_err\"]=0\n\n        if len(hess_uncerts)>0:\n          fit_outputs[idx][\"Tcenter_hess_err\"]=hess_uncerts[5]\n          fit_outputs[idx][\"T_hess_err\"]=hess_uncerts[6]\n          if len(ch0_hess_errs)==0:\n            fit_outputs[idx][\"fit_p3_hess_err\"]=hess_uncerts[1]\n            fit_outputs[idx][\"fit_p2_hess_err\"]=hess_uncerts[2]\n            fit_outputs[idx][\"fit_p1_hess_err\"]=hess_uncerts[3]\n            fit_outputs[idx][\"fit_p0_hess_err\"]=hess_uncerts[4]\n        else:\n          fit_outputs[idx][\"Tcenter_hess_err\"]=0\n          fit_outputs[idx][\"T_hess_err\"]=0\n\n\n        params=torch.tensor(central_poly.x,device=device)\n\n        fit_y,_=gpu_compute_poly_fit_func(params,(torch.tensor(ch0_signal_x_binned,device=device),Tcenter_guess,T_guess,tau_guess,MAX_LDC_COEF,torch.tensor(fgs_central.x[-4:],device=device)),include_signal=True,return_signal=False)\n        fit_y=torch.squeeze(fit_y,dim=0)\n    \n        fit_y_nosig,_=gpu_compute_poly_fit_func(params,(torch.tensor(ch0_signal_x_binned,device=device),Tcenter_guess,T_guess,tau_guess,MAX_LDC_COEF,torch.tensor(fgs_central.x[-4:],device=device)),include_signal=False,return_signal=False)\n        fit_y_nosig=torch.squeeze(fit_y_nosig,dim=0)\n        fit_y_sigonly=fit_y_nosig-fit_y\n\n        residual=(ch0_signal_y[0,i,:]-fit_y)/ch0_signal_yerr[0,i,:]\n\n        if pre_end>pre_start:\n          fit_outputs[idx][\"res_pre\"]=residual[pre_start:pre_end].mean()\n          fit_outputs[idx][\"bg_pre\"]=fit_y_nosig[pre_start:pre_end].mean()\n          fit_outputs[idx][\"sig_pre\"]=fit_y_sigonly[pre_start:pre_end].mean()\n        else:\n          fit_outputs[idx][\"res_pre\"]=0\n          fit_outputs[idx][\"bg_pre\"]=0\n          fit_outputs[idx][\"sig_pre\"]=0\n        if ingr_end>ingr_start:\n          fit_outputs[idx][\"res_ingr\"]=residual[ingr_start:ingr_end].mean()\n          fit_outputs[idx][\"bg_ingr\"]=fit_y_nosig[ingr_start:ingr_end].mean()\n          fit_outputs[idx][\"sig_ingr\"]=fit_y_sigonly[ingr_start:ingr_end].mean()\n        else:\n          fit_outputs[idx][\"res_ingr\"]=0\n          fit_outputs[idx][\"bg_ingr\"]=0\n          fit_outputs[idx][\"sig_ingr\"]=0\n        if mid_end>mid_start:\n          fit_outputs[idx][\"res_mid\"]=residual[mid_start:mid_end].mean()\n          fit_outputs[idx][\"bg_mid\"]=fit_y_nosig[mid_start:mid_end].mean()\n          fit_outputs[idx][\"sig_mid\"]=fit_y_sigonly[mid_start:mid_end].mean()\n        else:\n          fit_outputs[idx][\"res_mid\"]=0\n          fit_outputs[idx][\"bg_mid\"]=0\n          fit_outputs[idx][\"sig_mid\"]=0\n        if egr_end>egr_start:\n          fit_outputs[idx][\"res_egr\"]=residual[egr_start:egr_end].mean()\n          fit_outputs[idx][\"bg_egr\"]=fit_y_nosig[egr_start:egr_end].mean()\n          fit_outputs[idx][\"sig_egr\"]=fit_y_sigonly[egr_start:egr_end].mean()\n        else:\n          fit_outputs[idx][\"res_egr\"]=0\n          fit_outputs[idx][\"bg_egr\"]=0\n          fit_outputs[idx][\"sig_egr\"]=0\n        if residual.shape[-1]>post_start:\n          fit_outputs[idx][\"res_post\"]=residual[post_start:].mean()\n          fit_outputs[idx][\"bg_post\"]=fit_y_nosig[post_start:].mean()\n          fit_outputs[idx][\"sig_post\"]=fit_y_sigonly[post_start:].mean()\n        else:\n          fit_outputs[idx][\"res_post\"]=0\n          fit_outputs[idx][\"bg_post\"]=0\n          fit_outputs[idx][\"sig_post\"]=0\n\n\n      #for obs_name in [\"lincorr_fit_depth\"]:\n      #  print(\"planet \"+str(planet)+\" fgs \"+str(obs_name)+\"=\"+str(fit_outputs[0][obs_name]),flush=True)\n      #  print(\"planet \"+str(planet)+\" ch0 150 \"+str(obs_name)+\"=\"+str(fit_outputs[151][obs_name]),flush=True)\n\n      #  print(\"planet \"+str(planet)+\" fgs \"+str(obs_name)+\"_err=\"+str(fit_outputs[0][obs_name+\"_err\"]),flush=True)\n      #  print(\"planet \"+str(planet)+\" ch0 150 \"+str(obs_name)+\"_err=\"+str(fit_outputs[151][obs_name+\"_err\"]),flush=True)\n\n      #for obs_name in [\"raw_fit_depth\",\"Tcenter\",\"T\",\"tau\",\"ldc\",\"fit_p0\",\"fit_p1\",\"fit_p2\",\"fit_p3\"]:\n      #  print(\"planet \"+str(planet)+\" fgs \"+str(obs_name)+\"=\"+str(fit_outputs[0][obs_name]),flush=True)\n      #  print(\"planet \"+str(planet)+\" ch0 150 \"+str(obs_name)+\"=\"+str(fit_outputs[151][obs_name]),flush=True)\n\n      #  print(\"planet \"+str(planet)+\" fgs \"+str(obs_name)+\"_hess_err=\"+str(fit_outputs[0][obs_name+\"_hess_err\"]),flush=True)\n      #  print(\"planet \"+str(planet)+\" ch0 150 \"+str(obs_name)+\"_hess_err=\"+str(fit_outputs[151][obs_name+\"_hess_err\"]),flush=True)\n\n      names=[\"chisq_ndof\",\"sig_pre\",\"sig_ingr\",\"sig_mid\",\"sig_egr\",\"sig_post\"]\n      names+=[\"bg_pre\",\"bg_ingr\",\"bg_mid\",\"bg_egr\",\"bg_post\"]\n      names+=[\"res_pre\",\"res_ingr\",\"res_mid\",\"res_egr\",\"res_post\"]\n        \n      #for obs_name in names:\n      #  print(\"planet \"+str(planet)+\" fgs \"+str(obs_name)+\"=\"+str(fit_outputs[0][obs_name]),flush=True)\n      #  print(\"planet \"+str(planet)+\" ch0 150 \"+str(obs_name)+\"=\"+str(fit_outputs[151][obs_name]),flush=True)\n\n\n      postproc_start_time=time_lib.time()\n\n      #Now run the postprocessing networks\n      #print(\"about to run postprocessing network\",flush=True)\n      orbit_block=torch.tensor([list(star_params_norm[planet])],dtype=torch.float32,device=device)\n      #print(\"got orbit_block ok\",flush=True)\n      #postproc_in=[make_feature(o,var_max=nn_input_var_max,var_min=nn_input_var_min,orbit_block=orbit_block) for o in fit_outputs]\n      postproc_in=[]\n      for io,o in enumerate(fit_outputs):\n        #print(\"assembling features for wavelength \"+str(io),flush=True)\n        postproc_in.append(make_feature(o,var_max=nn_input_var_max,var_min=nn_input_var_min,orbit_block=orbit_block))\n\n        \n      #print(\"make_feature ok, type(postproc_in)=\"+str(type(postproc_in)),flush=True)\n      postproc_in=collate_example(postproc_in)\n      #print(\"collate ok, type(postproc_in)=\"+str(type(postproc_in)),flush=True)\n        \n      if postproc_models is None:\n        postproc_models,postproc_fallbacks=load_postproc_models(postproc_in)\n      #print(\"loaded models ok\",flush=True)\n      with torch.no_grad():\n        predictions=[m(postproc_in) for m in postproc_models]\n        fallbacks=[m(postproc_in) for m in postproc_fallbacks]\n      #print(\"got predictions ok\",flush=True)\n\n      #pull up some variables that will help us decide whether to go with the main prediction or the fallback.\n      #These selections have a small and not-always-positive effect in local cross-validation, but they\n      #do steer clear of places I worry about and distrust (failed/weirdo fits and the like)\n      outscore=postproc_in[\"poly_unnorm\"][:,:,1:4]\n      outscore=torch.sum(torch.abs(outscore),dim=-1)  #shape (batch_size,n_wavelengths)\n\n      #predictions should have shape (batch_size,2,n_wavelengths), and it's convenient if these can broadcast to that shape\n      outscore=torch.unsqueeze(outscore,dim=-2).expand(predictions[0].shape)\n\n      chisq=postproc_in[\"fit_unnorm\"][:,:,-1]  \n      chisq=torch.unsqueeze(chisq,dim=-2)\n\n      tau=postproc_in[\"sliding_window_in_unnorm\"][:,:,1]  \n      tau=torch.unsqueeze(tau,dim=-2)\n\n      tau_err=postproc_in[\"sliding_window_errs_in_unnorm\"][:,:,1]  #shape (batch_size,n_wavelengths)\n      tau_err=torch.unsqueeze(tau_err,dim=-2)\n\n      ldc=postproc_in[\"sliding_window_in_unnorm\"][:,:,2]\n      ldc=torch.unsqueeze(ldc,dim=-2)\n\n      res_pre=postproc_in['res_unnorm'][:,:,0]\n      res_pre=torch.unsqueeze(res_pre,dim=-2)\n\n      res_ingr=postproc_in['res_unnorm'][:,:,1]\n      res_ingr=torch.unsqueeze(res_ingr,dim=-2)\n\n      res_mid=postproc_in['res_unnorm'][:,:,2]\n      res_mid=torch.unsqueeze(res_mid,dim=-2)\n\n      res_egr=postproc_in['res_unnorm'][:,:,3]\n      res_egr=torch.unsqueeze(res_egr,dim=-2)\n\n      res_post=postproc_in['res_unnorm'][:,:,4]\n      res_post=torch.unsqueeze(res_post,dim=-2)\n\n      tstart_mean=postproc_in[\"extras_unnorm\"][:,:,0]\n      tstart_mean=torch.unsqueeze(tstart_mean,dim=-2)\n\n      tend_mean=postproc_in[\"extras_unnorm\"][:,:,1]\n      tend_mean=torch.unsqueeze(tend_mean,dim=-2)\n\n      tstart_sig=postproc_in[\"extras_unnorm\"][:,:,2]\n      tstart_sig=torch.unsqueeze(tstart_sig,dim=-2)\n\n      tend_sig=postproc_in[\"extras_unnorm\"][:,:,3]\n      tend_sig=torch.unsqueeze(tend_sig,dim=-2)\n\n      gressratio=torch.clip(postproc_in[\"extras_unnorm\"][:,:,2],1,5625)/torch.clip(postproc_in[\"extras_unnorm\"][:,:,3],1,5625)\n      gressratio=torch.where(gressratio>1,1/gressratio,gressratio)\n      gressratio=torch.unsqueeze(gressratio,dim=-2)\n\n      #don't use fallback if the initial (pre-fit) ingress/egress estimates were highly asymmetric; \n      #that is a sign that the initial ingress/egress finding failed, and that failure can propagate to the fallback.\n      safety1=torch.where(gressratio<0.5,1,0) \n\n      #don't use fallback if a transit starts or ends too close to either end of the time series\n      safety2=torch.where(tstart_mean-tstart_sig<0,1,0)  \n      safety3=torch.where(tend_mean+tend_sig>5625,1,0)\n\n      safety=torch.maximum(safety1,safety2)\n      safety=torch.maximum(safety,safety3)\n\n         \n      mask1=torch.where(torch.abs(outscore)<7,1,0)\n      mask2=torch.where(chisq<1.5,1,0)\n      mask3=torch.where(tau<40,1,0)\n      mask4=torch.where(torch.abs(res_pre)<1,1,0)\n      mask5=torch.where(torch.abs(res_ingr)<1,1,0)\n      mask6=torch.where(torch.abs(res_mid)<1,1,0)\n      mask7=torch.where(torch.abs(res_egr)<1,1,0)\n      mask8=torch.where(torch.abs(res_post)<1,1,0)\n      mask9=torch.where(tstart_mean>400,1,0)\n      mask10=torch.where(tend_mean<5625-400,1,0)\n      mask11=torch.where(ldc>0.1,1,0)\n      mask12=torch.where(tau_err<5,1,0)\n\n      mask=mask1*mask2*mask3*mask4*mask5*mask6*mask7*mask8*mask9*mask10*mask11*mask12\n\n      print(\"sum(mask) before safety: \"+str(torch.sum(mask))+\"; sum(safety)=\"+str(torch.sum(safety))+\"; sum(mask) after safety: \"+str(torch.sum(torch.maximum(mask,safety))),flush=True)\n      mask_list.append((torch.sum(mask),torch.sum(safety),torch.sum(torch.maximum(mask,safety))))   \n      mask=torch.maximum(mask,safety)\n      #print(\"mask=\"+str(mask),flush=True)\n      #print(\"nonzero elts=\"+str(torch.nonzero(mask)),flush=True)\n      #print(\"prediction 0 (before squash):\"+str(predictions[0]),flush=True)\n      #raise Exception(\"Stop\")\n        \n      #print(\"fallback 0:\"+str(fallbacks[0]),flush=True)\n\n      #use fallback prediction when conditions are triggered (was done for contest submission)\n      predictions=[torch.where(mask>0,pred,fb) for pred,fb in zip(predictions,fallbacks)]\n\n      #debugging:  get rid of fallbacks to rule out problems there\n      #predictions=predictions  #narf\n\n\n         \n      #the five models here are not really independent measurements, just different instances of the \n      #same netowrk trained on five different train/val splits.  I think a weighted average with the usual\n      #gaussian uncertainty propagation would not be appropriate.  Go with an arithmetic mean instead.\n      pred_avg=sum(predictions)/len(predictions)\n\n\n      #debugging:  check exception handling\n      #raise Exception(\"just testing\")\n\n         \n      #pred_avg should have shape (batch_size,2,n_wavelengths).  Elements in pred_avg[:,0,:] are predicted means,\n      #but elements in pred_avg[:,1,:] are predicted width \"logits\" which need to be passed through a softplus before interpretation\n      pred_avg[:,1,:]=torch.clip(torch.nn.functional.softplus(pred_avg[:,1,:]),1e-6,1)\n      predicted_spectra.append(pred_avg)\n      #print(\"done with pass through visit loop\",flush=True)\n\n      #print(\"pred_avg after squash=\"+str(pred_avg),flush=True)\n      load_time_list.append(cpu_fit_start_time-example_start_time)\n      cpu_fit_time_list.append(gpu_fit_start_time-cpu_fit_start_time)\n      gpu_fit_time_list.append(postproc_start_time-gpu_fit_start_time)\n      postproc_time_list.append(time_lib.time()-postproc_start_time)\n\n      print('this example took '+str(load_time_list[-1])+\" to load, \"+str(cpu_fit_time_list[-1])+\" for cpu fit, \"+str(gpu_fit_time_list[-1])+\" for gpu fit, and \"+str(postproc_time_list[-1])+\" for postprocessing\",flush=True)\n\n     except Exception as e:\n      print(\"\\n\\n!!!!--->>>caught an exception during the validation loop: \"+str(e),flush=True)\n      default_pred_mean=naive_mean*torch.ones((1,1,283),dtype=torch.float32)\n      default_pred_sigma=naive_sigma*torch.ones((1,1,283),dtype=torch.float32)\n      default_pred=torch.cat([default_pred_mean,default_pred_sigma],dim=1)\n      predicted_spectra.append(default_pred)\n      #raise e\n\n\n    #print(\"predicted_spectra=\"+str(predicted_spectra),flush=True)\n    #if we have only one observation of the transit, then the the one entry in predicted_spectra is the prediction\n    #But if we have more than one, we have independent measurements which should probably be combined using gaussian error propagation.\n    if len(predicted_spectra)==1:\n      pred=predicted_spectra[0]\n    else:\n      pred=combine_measurements(predicted_spectra)\n      #print(\"combined two visits!  Result: \"+str(pred),flush=True) \n\n    final_means.append(pred[0,0,:].detach().cpu().numpy())\n    final_sigmas.append(pred[0,1,:].detach().cpu().numpy())\n\n    #print(\"pred.shape=\"+str(pred.shape),flush=True)\n    #print(\"before clipping, new element in final means=\"+str(final_means[-1]),flush=True)\n    #print(\"...and sigmas=\"+str(final_sigmas[-1]),flush=True)\n      \n    #In case of nan, default to naive mean/sigma.  \n    #I expect/hope this condition pretty much never happens.\n    final_sigmas[-1]=np.where(np.isnan(final_means[-1]),naive_sigma,final_sigmas[-1].clip(0))\n    final_means[-1]=np.where(np.isnan(final_means[-1]),naive_mean,final_means[-1].clip(0))\n    final_sigmas[-1]=np.where(np.isinf(final_sigmas[-1]),naive_sigma,final_sigmas[-1].clip(0))\n    final_means[-1]=np.where(np.isinf(final_means[-1]),naive_mean,final_means[-1].clip(0))\n\n    final_means[-1]=np.expand_dims(final_means[-1],axis=0)\n    final_sigmas[-1]=np.expand_dims(final_sigmas[-1],axis=0)\n    #print(\"just appended to final_means and final_sigmas with shapes \"+str([final_means[-1].shape,final_sigmas[-1].shape]),flush=True)\n\n    #print(\"that includes means=\"+str(final_means[-1]),flush=True)\n    #print(\"...and sigmas=\"+str(final_sigmas[-1]),flush=True)\n  print(\"average run time: \"+str(sum(run_time_list)/len(run_time_list)),flush=True)\n  print(\"average load time: \"+str(sum(load_time_list)/len(load_time_list)),flush=True)\n  print(\"average cpu fit time: \"+str(sum(cpu_fit_time_list)/len(cpu_fit_time_list)),flush=True)\n  print(\"average gpu fit time: \"+str(sum(gpu_fit_time_list)/len(gpu_fit_time_list)),flush=True)\n  print(\"average postproc time: \"+str(sum(postproc_time_list)/len(postproc_time_list)),flush=True)\n  print('average mask (pre-safety): '+str(sum([tup[0] for tup in mask_list])/len(mask_list)),flush=True)\n  print('average safety: '+str(sum([tup[1] for tup in mask_list])/len(mask_list)),flush=True)\n  print('average mask (post-safety): '+str(sum([tup[2] for tup in mask_list])/len(mask_list)),flush=True)\n    \nexcept Exception as e:\n  print(\"caught an exception: \"+str(e),flush=True)\n  run_is_ok=False\n  raise e\n\n#based on the 2024 competition, expect the following possible outputs from a submission:\n#  - a valid score, which may or may not be zero\n#  - a \"notebook timeout\" error if it runs too long\n#  - a \"submission csv not found\" if the run finishes in time but does not produce a submission.csv\n#  - a \"submission scoring error\" if there is a submission.csv but it contains a formatting problem or something\n#  - a \"notebook threw exception\" error if there is an uncaught exception\ntry:    \n  ss = pd.read_csv('/kaggle/input/ariel-data-challenge-2025/sample_submission.csv')\n\n  final_means=np.concatenate(final_means,axis=0)\n  final_sigmas=np.concatenate(final_sigmas,axis=0)\n  print(\"run is ok, final_means shape=\"+str(final_means.shape))\n  print(\"...and final_sigmas shape=\"+str(final_sigmas.shape))\n\n  final_means=np.where(np.isnan(final_means),naive_mean,final_means)\n  final_means=np.where(np.isinf(final_means),naive_mean,final_means)\n  final_sigmas=np.where(np.isnan(final_sigmas),naive_sigma,final_sigmas)\n  final_sigmas=np.where(np.isinf(final_sigmas),naive_sigma,final_sigmas)\n\n  planet_id_np=np.array([[planet] for planet in meta_dict['index']])\n  print(\"planet_id_np.shape=\"+str(planet_id_np.shape),flush=True)\n  print(\"final_means.shape=\"+str(final_means.shape),flush=True)\n  print(\"final_sigmas.shape=\"+str(final_sigmas.shape),flush=True)\n    \n  submission = pd.DataFrame(np.concatenate([planet_id_np,final_means,final_sigmas], axis=1), columns=ss.columns)\n  #submission['planet_id'] = meta_dict[\"index\"]\n  submission=submission.set_index('planet_id')\n  \n  submission.to_csv('submission.csv')\n\n  #print(\"submission=\"+str(submission),flush=True)\n  #print(\"final means: \"+str(final_means),flush=True)\n  #print(\"final sigmas: \"+str(final_sigmas),flush=True)\n\n  #with open('submission.csv','r') as f:\n  #    for line in f:\n  #      print(line)\n          \nexcept Exception as e:\n  print(\"exception trying to build output -- do not make a submission! Exception was: \"+str(e))\n    \n    \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-09T03:06:17.151971Z","iopub.execute_input":"2025-10-09T03:06:17.152203Z","iopub.status.idle":"2025-10-09T03:07:14.05876Z","shell.execute_reply.started":"2025-10-09T03:06:17.152182Z","shell.execute_reply":"2025-10-09T03:07:14.058001Z"}},"outputs":[],"execution_count":null}]}