#!/usr/bin/env python3
"""CPU-only geometry and input-buffer estimator, not an HSI model benchmark.
Python 3.9+, standard library only. Never downloads data or starts model jobs.
Examples:
  python toy_estimator.py
  python toy_estimator.py --mode native --inspect mamba --axis spatial
  python toy_estimator.py --protocol hsi-comparison-protocol.json
"""
import argparse,json,math,pathlib
DEFAULTS={'mode':'matched','inspect':'cnn','split':'scene','target':'hidden','pretrain':'none','scaling':'train','side':128,'bands':48,'patch':25,'batch':16,'dtype':'fp32','trials':8,'seeds':5,'minutes':60,'axis':'spectral','heads':4,'dim':64,'state':16}
ENUMS={'mode':['matched','native'],'inspect':['cnn','transformer','mamba'],'split':['scene','block','random'],'target':['hidden','visible'],'pretrain':['none','equal','mixed'],'scaling':['train','all'],'dtype':['fp32','fp16'],'axis':['spectral','spatial']}
BOUNDS={'side':(32,1024),'bands':(8,256),'patch':(3,31),'batch':(1,64),'trials':(1,50),'seeds':(1,20),'minutes':(1,240),'heads':(1,16),'dim':(8,256),'state':(1,64)}
def normalize(raw):
 s=dict(DEFAULTS)
 for k,choices in ENUMS.items():
  if raw.get(k) in choices:s[k]=raw[k]
 for k,(lo,hi)in BOUNDS.items():
  try:
   v=float(raw.get(k,''))
   if math.isfinite(v):s[k]=min(hi,max(lo,math.floor(v+0.5)))
  except (ValueError,TypeError):pass
 if s['patch']%2==0:s['patch']=min(31,s['patch']+1)
 return s

def footprint(size,center=(32,24),side=64):
 if size=='whole':return {(y,x)for y in range(side)for x in range(side)}
 r=(size-1)//2
 return {(y,x)for y in range(max(0,center[0]-r),min(side,center[0]+r+1))for x in range(max(0,center[1]-r),min(side,center[1]+r+1))}

def overlap(size):
 a,b=footprint(size),footprint(size,(32,32));i=len(a&b)
 return {'first':len(a),'second':len(b),'intersection':i,'union':len(a|b),'fraction':i/len(a)}

def compute(raw=None):
 s=normalize(raw or {});n=s['side']**2;b=4 if s['dtype']=='fp32'else 2;formats=[]
 for id,name in [('cnn','CNN / HybridSN-inspired'),('transformer','SpectralFormer-inspired'),('mamba','MambaHSI-inspired')]:
  dense=s['mode']=='native'and id=='mamba';p=s['patch']if s['mode']=='matched'else(25 if id=='cnn'else 7);support=n if dense else p*p;batch=1 if dense else s['batch']
  formats.append({'id':id,'name':name,'dense':dense,'context':'whole scene'if dense else f'{p} × {p} patch','patch':None if dense else p,'inputValuesPerSample':support*s['bands'],'inputBytesPerCall':batch*support*s['bands']*b,'predictionTargetsPerCall':n if dense else s['batch'],'callSamples':batch,'naiveDenseMapInputValues':n*s['bands']*(1 if dense else p*p),'overlap':overlap('whole'if dense else p)})
 selected=next(m for m in formats if m['id']==s['inspect']);L=s['bands']+1 if s['axis']=='spectral'else(n if selected['dense']else selected['patch']**2)
 return {'settings':s,'formats':formats,'generic':{'tokens':L,'axis':s['axis'],'bytesPerScalar':b,'materializedScoresBytes':s['heads']*L*L*b,'oneTokenFeatureBytes':L*s['dim']*b,'oneStreamingStateBytes':s['dim']*s['state']*b},'budget':{'trialCountPerModel':s['trials'],'searchSeedsPerTrial':1,'finalSeeds':s['seeds'],'perRunCapMinutes':s['minutes'],'upperBoundRunsThreeModels':3*(s['trials']+s['seeds']),'upperBoundDeviceHoursThreeModels':3*(s['trials']+s['seeds'])*s['minutes']/60},'measurements':{'accuracy':None,'latency':None,'peakMemory':None},'status':'unexecuted-protocol-and-exact-toy-accounting','disclaimer':'Exact input and component arithmetic only. No trained models, benchmark accuracy, FLOPs, latency or total-memory estimates. Footprint grid is 64 by 64, separate from the scene-size control.'}

def main():
 p=argparse.ArgumentParser(description=__doc__,formatter_class=argparse.RawDescriptionHelpFormatter);p.add_argument('--protocol',type=pathlib.Path);p.add_argument('--output',type=pathlib.Path)
 for k in DEFAULTS:p.add_argument('--'+k,choices=ENUMS.get(k),type=int if k in BOUNDS else str)
 a=vars(p.parse_args());source=a.pop('protocol');output=a.pop('output');settings={}
 if source:
  obj=json.loads(source.read_text());settings=obj.get('settings',obj)
  if not isinstance(settings,dict):p.error('Protocol settings must be an object')
 settings.update({k:v for k,v in a.items()if v is not None});result=json.dumps(compute(settings),indent=2,ensure_ascii=False)+'\n'
 if output:output.write_text(result)
 else:print(result,end='')
if __name__=='__main__':main()
