#!/usr/bin/env python """ Aether Phase-1 PRE-EXTRACT: run the frozen encoders ONCE, cache their features to disk. This is the portability enabler — after this, alignment/adapters train on the cached tensors with NO encoders loaded (fits 6900XT / Kaggle-T4 / Colab-T4 / any small GPU). Writes /work/slivers_cache//.pt (raw encoder feature, pre-projection) + /work/slivers_cache/index.jsonl ({id, modality, caption, feat_path, shape}) Encoders live only here (the big box); the cache travels. Run on the rental MI300 (has all 3 encoders). Validated on a tiny mixed set. """ import os, sys, json, torch sys.path.insert(0,"/work") import forward as F dev=F.dev CACHE="/work/slivers_cache" def main(): torch.set_grad_enabled(False) m=F.AetherPhase1().to(dev); m.eval() os.makedirs(CACHE, exist_ok=True) idx=[] torch.manual_seed(42) # ---- VISION items (image -> caption) ---- os.makedirs(f"{CACHE}/vision", exist_ok=True) vcaps=["a red brick electrical panel on a wall","a blue ceramic coffee mug on a desk", "three yellow pencils in a glass jar","a green circuit board with silver traces"] for i,cap in enumerate(vcaps): px=torch.randn(1,3,448,448) # (real pipeline: load actual Tier-A image) feat=m.enc_vision(px).cpu() # (1025,3200) frozen features p=f"{CACHE}/vision/v{i}.pt"; torch.save(feat,p) idx.append({"id":f"v{i}","modality":"vision","caption":cap,"feat":p,"shape":list(feat.shape)}) # ---- AUDIO items (speech -> transcript) ---- os.makedirs(f"{CACHE}/audio", exist_ok=True) acaps=["the meeting starts at nine tomorrow","turn off the main breaker first"] for i,cap in enumerate(acaps): mel=torch.randn(400,128); alen=torch.tensor([400]) feat=m.enc_audio(mel,alen).cpu() # (~100,1280) p=f"{CACHE}/audio/a{i}.pt"; torch.save(feat,p) idx.append({"id":f"a{i}","modality":"audio","caption":cap,"feat":p,"shape":list(feat.shape)}) # ---- 3D geometry items (mesh -> caption); pre-extract SLAT latent + voxel coords ---- os.makedirs(f"{CACHE}/geom", exist_ok=True) gcaps=["a low-poly wooden crate","a cylindrical steel conduit fitting"] for i,cap in enumerate(gcaps): M=300; vz=torch.randint(0,64,(M,3)); b=torch.zeros(M,1) coords=torch.cat([b,vz.flip(-1).float()],1); feats=torch.randn(M,1024) gf,gc=m.enc_geom(coords,feats) # (M',8),(M',4) gcz=(gc[:,[1,2,3]].float()/gc[:,[1,2,3]].float().max().clamp(min=1)*15).long() p=f"{CACHE}/geom/g{i}.pt"; torch.save({"feat":gf.cpu(),"xyz":gcz.cpu()},p) idx.append({"id":f"g{i}","modality":"geom","caption":cap,"feat":p,"shape":list(gf.shape)}) json.dump(idx, open(f"{CACHE}/index.jsonl","w"), indent=2) print(f"PRE-EXTRACT OK: {len(idx)} items cached to {CACHE}") for r in idx: print(" ",r["id"],r["modality"],r["shape"],"->",os.path.getsize(r["feat"])//1024,"KB") if __name__=="__main__": main()