import torch, json
from pathlib import Path
from safetensors import safe_open
from transformers import SiglipVisionConfig, SiglipVisionModel
from huggingface_hub import hf_hub_download
import numpy as np
root=Path('artifacts/pretrained-siglip-20260912/fixture')
torch.set_grad_enabled(False); torch.set_num_threads(8)
p=hf_hub_download('google/siglip-so400m-patch14-384','model.safetensors',revision='9fdffc58afc957d1a03a25b10dba0329ab15c2a3')
c=SiglipVisionConfig.from_pretrained('google/siglip-so400m-patch14-384',revision='9fdffc58afc957d1a03a25b10dba0329ab15c2a3'); c._attn_implementation='sdpa'
with torch.device('meta'): m=SiglipVisionModel(c)
with safe_open(p,framework='pt') as f: m.load_state_dict({k.removeprefix('vision_model.'):f.get_tensor(k) for k in f.keys() if k.startswith('vision_model.')},assign=True)
m.embeddings.position_ids=torch.arange(729).expand(1,-1)
m.eval().to('cuda')
a=np.fromfile(root/'authors/patches.f32',dtype='<f4').reshape(1,27,27,14,14,3)
pix=torch.from_numpy(a.transpose(0,5,1,3,2,4).reshape(1,3,378,378).copy()).cuda(); pix=torch.nn.functional.pad(pix,(0,6,0,6))
stats={}
for name,module in m.named_modules():
 if isinstance(module,(torch.nn.Linear,torch.nn.Conv2d)):
  def hook(module,args,name=name):
   x=args[0]; stats[name]={'max':x.abs().max().item(),'over15':(x.abs()>15).float().mean().item(),'weightMax':module.weight.abs().max().item()}
  module.register_forward_pre_hook(hook)
ref=m(pix).pooler_output
expected=torch.from_numpy(np.fromfile(root/'authors/expected.f32',dtype='<f4')).cuda()
print('cuda-vs-cpu',torch.nn.functional.cosine_similarity(ref.flatten(),expected,dim=0).item(),flush=True)
print(json.dumps(stats,indent=2),flush=True)
def quant(x):
 return ((x*16).clamp(-240,240).to(torch.float8_e4m3fnuz).float()/16).to(x.dtype)
for name,module in m.named_modules():
 if isinstance(module,(torch.nn.Linear,torch.nn.Conv2d)):
  module.weight.copy_(quant(module.weight))
  module.register_forward_pre_hook(lambda module,args:(quant(args[0]),)+args[1:])
 if isinstance(module,torch.nn.MultiheadAttention):
  module.in_proj_weight.copy_(quant(module.in_proj_weight))
  module.register_forward_pre_hook(lambda module,args:tuple(quant(x) for x in args))
# This is an optimistic FP8-operands/FP32-accumulation probe, not an IPU simulator.
out=m(pix).pooler_output
print('FP8 operands fixed -4, FP32 accumulation cosine', torch.nn.functional.cosine_similarity(out.flatten(),expected,dim=0).item(),flush=True)
