""" Training Script for Φ+G Pipeline ================================= Run: python train.py --phase 1 --data_dir data/ --output_dir checkpoints/ Requires: pip install torch mujoco numpy Data: generate with generate_data.py first. Phases: 1: Encoder (supervised) 2: Vectorizer (contrastive) 3a: Cross-encoder alignment (contrastive) 3b: Cross-encoder action (behavioral cloning) """ import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import Dataset, DataLoader import numpy as np import json import os import argparse from pathlib import Path # Import models from pipeline.py (should be in same directory or on PYTHONPATH) # If running standalone, the models are defined inline below. def filter_manifest_entries(data_dir, manifest, *keys): kept = [] skipped = 0 for entry in manifest: if all((data_dir / entry[key]).exists() for key in keys): kept.append(entry) else: skipped += 1 if skipped: print(f"[train.py] Skipping {skipped} manifest entries with missing files in {data_dir}") return kept # ============================================================ # MODELS (self-contained — same as pipeline.py) # ============================================================ class EncodingModel(nn.Module): def __init__(self, latent_dim=64, max_objects=6): super().__init__() self.latent_dim = latent_dim self.max_objects = max_objects sd = 64 self.backbone = nn.Sequential( nn.Conv2d(3,16,7,2,3),nn.ReLU(),nn.Conv2d(16,32,5,2,2),nn.ReLU(), nn.Conv2d(32,64,3,2,1),nn.ReLU(),nn.Conv2d(64,128,3,2,1),nn.ReLU()) self.slots_init = nn.Parameter(torch.randn(1,max_objects,sd)*0.02) self.sq=nn.Linear(sd,32);self.sk=nn.Linear(128,32);self.sv=nn.Linear(128,sd) self.sgru=nn.GRUCell(sd,sd) self.smlp=nn.Sequential(nn.Linear(sd,sd),nn.ReLU(),nn.Linear(sd,sd)) self.to_z=nn.Linear(sd,latent_dim) self.to_pose=nn.Linear(sd,7) self.to_exist=nn.Linear(sd,1) self.edge_net=nn.Sequential(nn.Linear(sd*2,32),nn.ReLU(),nn.Linear(32,33)) def forward(self, image): B,N=image.shape[0],self.max_objects f=self.backbone(image).flatten(2).permute(0,2,1) slots=self.slots_init.expand(B,-1,-1) for _ in range(3): q,k,v=self.sq(slots),self.sk(f),self.sv(f) a=F.softmax(torch.einsum('bnd,bmd->bnm',q,k)/6,dim=1) a=a/(a.sum(-1,keepdim=True)+1e-8) u=torch.einsum('bnm,bmd->bnd',a,v) slots=self.sgru(u.reshape(B*N,-1),slots.reshape(B*N,-1)).reshape(B,N,-1) slots=slots+self.smlp(slots) z=self.to_z(slots) poses=self.to_pose(slots) exist=torch.sigmoid(self.to_exist(slots).squeeze(-1)) si=slots.unsqueeze(2).expand(-1,-1,N,-1) sj=slots.unsqueeze(1).expand(-1,N,-1,-1) eo=self.edge_net(torch.cat([si,sj],-1)) return {'phi':{'z':z,'poses':poses,'existence':exist}, 'graph':{'edge_features':eo[...,:32],'contact_probs':torch.sigmoid(eo[...,32])}} class GNNLayer(nn.Module): def __init__(self, hd, ed): super().__init__() self.msg=nn.Sequential(nn.Linear(hd*2+ed,hd),nn.ReLU(),nn.Linear(hd,hd)) self.upd=nn.Sequential(nn.Linear(hd*2,hd),nn.ReLU(),nn.Linear(hd,hd)) self.norm=nn.LayerNorm(hd) def forward(self,h,ef,cp,mask): B,N,D=h.shape hi=h.unsqueeze(2).expand(-1,-1,N,-1) hj=h.unsqueeze(1).expand(-1,N,-1,-1) m=self.msg(torch.cat([hi,hj,ef],-1))*cp.unsqueeze(-1)*mask.unsqueeze(1) return self.norm(h+self.upd(torch.cat([h,m.sum(2)],-1)))*mask class StateVectorizer(nn.Module): def __init__(self, latent_dim=64, output_dim=256): super().__init__() hd=128 self.node_emb=nn.Sequential(nn.Linear(latent_dim+8,hd),nn.ReLU(),nn.Linear(hd,hd)) self.gnns=nn.ModuleList([GNNLayer(hd,32) for _ in range(2)]) self.aq=nn.Linear(hd,32);self.ak=nn.Linear(hd,32) self.proj=nn.Sequential(nn.Linear(hd,hd),nn.ReLU(),nn.Linear(hd,output_dim)) self.norm=nn.LayerNorm(output_dim) def forward(self,phi,graph): B,N=phi['z'].shape[:2] h=self.node_emb(torch.cat([phi['z'],phi['poses'],phi['existence'].unsqueeze(-1)],-1)) mask=phi['existence'].unsqueeze(-1);h=h*mask for g in self.gnns: h=g(h,graph['edge_features'],graph['contact_probs'],mask) q=self.aq(h.mean(1,keepdim=True));k=self.ak(h) a=F.softmax(torch.einsum('bid,bjd->bij',q,k)/6+(1-mask.transpose(1,2))*-1e9,-1) return self.norm(self.proj(torch.einsum('bij,bjd->bid',a,h).squeeze(1))) class CrossEncoder(nn.Module): def __init__(self, state_dim=256, embed_dim=256, action_dim=6): super().__init__() self.scene_proj=nn.Sequential(nn.Linear(state_dim,embed_dim),nn.ReLU(),nn.Linear(embed_dim,embed_dim)) self.scene_norm=nn.LayerNorm(embed_dim) self.text_emb=nn.Embedding(10000,128) self.text_pos=nn.Parameter(torch.randn(1,64,128)*0.02) self.text_tf=nn.TransformerEncoder( nn.TransformerEncoderLayer(128,4,256,batch_first=True,dropout=0.1),num_layers=2) self.text_proj=nn.Linear(128,embed_dim);self.text_norm=nn.LayerNorm(embed_dim) self.cross_attn=nn.MultiheadAttention(embed_dim,4,batch_first=True) self.fuse_norm=nn.LayerNorm(embed_dim) self.fuse_mlp=nn.Sequential(nn.Linear(embed_dim,embed_dim),nn.ReLU(),nn.Linear(embed_dim,embed_dim)) self.action_head=nn.Sequential(nn.Linear(embed_dim,128),nn.ReLU(),nn.Linear(128,action_dim),nn.Tanh()) self.logit_scale=nn.Parameter(torch.tensor(1/0.07).log()) def encode_scene(self,s): return self.scene_norm(self.scene_proj(s)) def encode_text(self,tok,mask=None): x=self.text_emb(tok)+self.text_pos[:,:tok.shape[1]] return self.text_norm(self.text_proj(self.text_tf(x,src_key_padding_mask=mask)[:,0])) def forward(self,s,text_tokens=None,text_mask=None): se=self.encode_scene(s) te=self.encode_text(text_tokens,text_mask) if text_tokens is not None else None toks=[se.unsqueeze(1)] if te is not None: toks.append(te.unsqueeze(1)) seq=torch.cat(toks,1) f,_=self.cross_attn(seq,seq,seq) fused=self.fuse_norm(f[:,0]+self.fuse_mlp(f[:,0])) return {'action':self.action_head(fused),'scene_emb':se,'text_emb':te,'fused':fused} def contrastive_loss(self,se,te): s,t=F.normalize(se,-1),F.normalize(te,-1) l=self.logit_scale.exp().clamp(max=100)*s@t.T lb=torch.arange(len(l),device=l.device) return(F.cross_entropy(l,lb)+F.cross_entropy(l.T,lb))/2 # ============================================================ # DATASETS # ============================================================ class Phase1Dataset(Dataset): """Supervised encoder training data.""" def __init__(self, data_dir, max_objects=6): self.data_dir = Path(data_dir) with open(self.data_dir / "manifest.json") as f: self.manifest = filter_manifest_entries(self.data_dir, json.load(f), "file") self.max_objects = max_objects def __len__(self): return len(self.manifest) def _pad(self, arr, target_n, fill=0): """Pad array's first dim to target_n.""" n = arr.shape[0] if n >= target_n: return arr[:target_n] pad_shape = (target_n - n,) + arr.shape[1:] return np.concatenate([arr, np.full(pad_shape, fill, dtype=arr.dtype)]) def __getitem__(self, idx): entry = self.manifest[idx] data = np.load(self.data_dir / entry["file"]) N = self.max_objects n_obj = int(data["n_objects"]) # Pad to max_objects poses = self._pad(data["gt_poses"], N) existence = self._pad(data["gt_existence"], N) materials = self._pad(data["gt_materials"], N) contacts = np.zeros((N, N), dtype=np.float32) c = data["gt_contacts"] contacts[:c.shape[0], :c.shape[1]] = c # Image: [H,W,3] uint8 → [3,H,W] float32 normalized image = data["image"].astype(np.float32).transpose(2, 0, 1) / 255.0 return { "image": torch.from_numpy(image), "gt_poses": torch.from_numpy(poses), "gt_existence": torch.from_numpy(existence), "gt_contacts": torch.from_numpy(contacts), "gt_materials": torch.from_numpy(materials), "sdf_samples": torch.from_numpy(data["sdf_samples"]), "n_objects": n_obj, } class Phase2Dataset(Dataset): """Contrastive scene pairs for vectorizer.""" def __init__(self, data_dir, max_objects=6): self.data_dir = Path(data_dir) with open(self.data_dir / "manifest.json") as f: self.manifest = filter_manifest_entries( self.data_dir, json.load(f), "positive_file", "negative_file" ) self.max_objects = max_objects def __len__(self): return len(self.manifest) def _pad(self, arr, n, fill=0): if arr.shape[0] >= n: return arr[:n] return np.concatenate([arr, np.full((n-arr.shape[0],)+arr.shape[1:], fill, dtype=arr.dtype)]) def __getitem__(self, idx): entry = self.manifest[idx] pos = np.load(self.data_dir / entry["positive_file"]) neg = np.load(self.data_dir / entry["negative_file"]) N = self.max_objects def pack(data): return { "gt_poses": torch.from_numpy(self._pad(data["gt_poses"], N)), "gt_existence": torch.from_numpy(self._pad(data["gt_existence"], N)), "gt_contacts": torch.from_numpy(self._pad( np.pad(data["gt_contacts"], ((0,max(0,N-data["gt_contacts"].shape[0])), (0,max(0,N-data["gt_contacts"].shape[1]))), mode='constant')[:N,:N], N)[:N,:N] if data["gt_contacts"].shape[0] < N else data["gt_contacts"][:N,:N]), "gt_materials": torch.from_numpy(self._pad(data["gt_materials"], N)), } # Simpler: just pad contacts directly def pack_v2(data): n = int(data["n_objects"]) poses = self._pad(data["gt_poses"], N) exist = self._pad(data["gt_existence"], N) mats = self._pad(data["gt_materials"], N) contacts = np.zeros((N, N), dtype=np.float32) c = data["gt_contacts"] contacts[:c.shape[0], :c.shape[1]] = c return { "gt_poses": torch.from_numpy(poses), "gt_existence": torch.from_numpy(exist), "gt_contacts": torch.from_numpy(contacts), "gt_materials": torch.from_numpy(mats), } return {"positive": pack_v2(pos), "negative": pack_v2(neg)} class Phase3aDataset(Dataset): """Scene-text pairs for contrastive alignment.""" def __init__(self, data_dir, max_objects=6, text_style="detailed"): self.data_dir = Path(data_dir) with open(self.data_dir / "manifest.json") as f: self.manifest = filter_manifest_entries(self.data_dir, json.load(f), "file") self.max_objects = max_objects self.text_style = text_style def __len__(self): return len(self.manifest) def __getitem__(self, idx): entry = self.manifest[idx] with open(self.data_dir / entry["file"]) as f: data = json.load(f) N = self.max_objects n = data["n_objects"] poses = np.array(data["gt_poses"], dtype=np.float32) exist = np.array(data["gt_existence"], dtype=np.float32) mats = np.array(data["gt_materials"], dtype=np.float32) contacts = np.array(data["gt_contacts"], dtype=np.float32) # Pad def pad1(a, n_target): if len(a) >= n_target: return a[:n_target] return np.concatenate([a, np.zeros((n_target-len(a),)+a.shape[1:], dtype=a.dtype)]) poses = pad1(poses, N) exist = np.concatenate([exist, np.zeros(max(0, N-len(exist)))])[:N] mats = pad1(mats, N) c_pad = np.zeros((N, N), dtype=np.float32) c_pad[:contacts.shape[0], :contacts.shape[1]] = contacts tok = data["tokenized"][self.text_style] return { "gt_poses": torch.from_numpy(poses), "gt_existence": torch.from_numpy(exist.astype(np.float32)), "gt_contacts": torch.from_numpy(c_pad), "gt_materials": torch.from_numpy(mats), "token_ids": torch.tensor(tok["token_ids"], dtype=torch.long), "attention_mask": torch.tensor(tok["attention_mask"], dtype=torch.bool), } class Phase3bDataset(Dataset): """Demonstration data for behavioral cloning.""" def __init__(self, data_dir, max_objects=6): self.data_dir = Path(data_dir) with open(self.data_dir / "manifest.json") as f: self.manifest = filter_manifest_entries(self.data_dir, json.load(f), "file") self.max_objects = max_objects def __len__(self): # Each demo has multiple timesteps return sum(e["n_steps"] for e in self.manifest) def __getitem__(self, idx): # Find which demo and which timestep cumulative = 0 for entry in self.manifest: if idx < cumulative + entry["n_steps"]: t = idx - cumulative break cumulative += entry["n_steps"] else: entry = self.manifest[-1] t = 0 data = np.load(self.data_dir / entry["file"]) N = self.max_objects def pad1(a, n_target): if len(a) >= n_target: return a[:n_target] return np.concatenate([a, np.zeros((n_target-len(a),)+a.shape[1:], dtype=a.dtype)]) poses = pad1(data["gt_poses"][t], N) exist = np.concatenate([data["gt_existence"], np.zeros(max(0, N-len(data["gt_existence"])))])[:N] mats = pad1(data["gt_materials"], N) contacts = np.zeros((N, N), dtype=np.float32) c = data["gt_contacts"][t] contacts[:c.shape[0], :c.shape[1]] = c return { "gt_poses": torch.from_numpy(poses.astype(np.float32)), "gt_existence": torch.from_numpy(exist.astype(np.float32)), "gt_contacts": torch.from_numpy(contacts), "gt_materials": torch.from_numpy(mats.astype(np.float32)), "action": torch.from_numpy(data["actions"][t].astype(np.float32)), "token_ids": torch.from_numpy(data["token_ids"].astype(np.int64)), "attention_mask": torch.from_numpy(data["attention_mask"]), } # ============================================================ # TRAINING LOOPS # ============================================================ def train_phase1(data_dir, output_dir, epochs=50, lr=3e-4, batch_size=8): """Phase 1: Train encoder (supervised).""" print("=== Phase 1: Encoder Training (Supervised) ===") os.makedirs(output_dir, exist_ok=True) dataset = Phase1Dataset(data_dir, max_objects=6) loader = DataLoader(dataset, batch_size=batch_size, shuffle=True, drop_last=True) model = EncodingModel(latent_dim=64, max_objects=6) optimizer = torch.optim.Adam(model.parameters(), lr=lr) log = {"epoch": [], "loss": [], "pose_err": [], "exist_err": [], "contact_err": []} for epoch in range(epochs): epoch_loss = 0 epoch_pose_err = 0 epoch_exist_err = 0 epoch_contact_err = 0 n_batches = 0 for batch in loader: out = model(batch["image"]) # Pose loss (MSE on position, ignoring quaternion for simplicity) pose_loss = F.mse_loss(out['phi']['poses'][:, :, :3], batch["gt_poses"][:, :, :3]) # Existence loss exist_loss = F.binary_cross_entropy(out['phi']['existence'], batch["gt_existence"]) # Contact loss contact_loss = F.binary_cross_entropy(out['graph']['contact_probs'], batch["gt_contacts"]) loss = pose_loss + exist_loss + 0.5 * contact_loss optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() epoch_loss += loss.item() epoch_pose_err += pose_loss.item() epoch_exist_err += exist_loss.item() epoch_contact_err += contact_loss.item() n_batches += 1 avg = lambda x: x / max(n_batches, 1) log["epoch"].append(epoch) log["loss"].append(avg(epoch_loss)) log["pose_err"].append(avg(epoch_pose_err)) log["exist_err"].append(avg(epoch_exist_err)) log["contact_err"].append(avg(epoch_contact_err)) if epoch % 10 == 0 or epoch == epochs - 1: print(f" Epoch {epoch:4d} | loss={avg(epoch_loss):.4f} " f"pose={avg(epoch_pose_err):.4f} exist={avg(epoch_exist_err):.4f} " f"contact={avg(epoch_contact_err):.4f}") # Save torch.save(model.state_dict(), os.path.join(output_dir, "encoder.pt")) with open(os.path.join(output_dir, "phase1_log.json"), "w") as f: json.dump(log, f) print(f" Saved: {output_dir}/encoder.pt") return model, log def train_phase2(data_dir, output_dir, encoder, epochs=50, lr=3e-4, batch_size=8): """Phase 2: Train vectorizer (contrastive).""" print("\n=== Phase 2: Vectorizer Training (Contrastive) ===") os.makedirs(output_dir, exist_ok=True) dataset = Phase2Dataset(data_dir, max_objects=6) loader = DataLoader(dataset, batch_size=batch_size, shuffle=True, drop_last=True) vectorizer = StateVectorizer(latent_dim=64, output_dim=256) optimizer = torch.optim.Adam(vectorizer.parameters(), lr=lr) temperature = 0.07 encoder.eval() log = {"epoch": [], "loss": [], "accuracy": []} for epoch in range(epochs): epoch_loss = 0 epoch_acc = 0 n_batches = 0 for batch in loader: B = batch["positive"]["gt_poses"].shape[0] # Create synthetic Φ+G from GT (bypass encoder for phase 2) # In production, would run encoder on two different camera views phi_pos = { 'z': torch.randn(B, 6, 64) * 0.1, # placeholder — in production, from encoder 'poses': batch["positive"]["gt_poses"], 'existence': batch["positive"]["gt_existence"], } graph_pos = { 'edge_features': torch.randn(B, 6, 6, 32) * 0.01, 'contact_probs': batch["positive"]["gt_contacts"], } # Positive: same scene, "different view" = add small noise to z phi_pos2 = { 'z': phi_pos['z'] + torch.randn_like(phi_pos['z']) * 0.01, 'poses': phi_pos['poses'] + torch.randn_like(phi_pos['poses']) * 0.001, 'existence': phi_pos['existence'], } graph_pos2 = graph_pos # same graph # Negative: different materials phi_neg = { 'z': torch.randn(B, 6, 64) * 0.1, # different z 'poses': batch["negative"]["gt_poses"], 'existence': batch["negative"]["gt_existence"], } # Encode material info into z for the negative # (In production, z comes from encoder which sees the actual materials) mat_pos = batch["positive"]["gt_materials"] # [B, N, 4] mat_neg = batch["negative"]["gt_materials"] # Inject material signal into z so contrastive can distinguish phi_pos['z'][:, :, :4] = mat_pos[:, :6, :] # first 4 dims of z = materials phi_pos2['z'][:, :, :4] = mat_pos[:, :6, :] + torch.randn(B, 6, 4) * 0.005 phi_neg['z'][:, :, :4] = mat_neg[:, :6, :] s_a = vectorizer(phi_pos, graph_pos) # [B, 256] s_b = vectorizer(phi_pos2, graph_pos2) # [B, 256] — positive pair # InfoNCE loss s_a_norm = F.normalize(s_a, dim=-1) s_b_norm = F.normalize(s_b, dim=-1) logits = s_a_norm @ s_b_norm.T / temperature # [B, B] labels = torch.arange(B) loss = (F.cross_entropy(logits, labels) + F.cross_entropy(logits.T, labels)) / 2 # Accuracy: positive pair should be closest acc = (logits.argmax(dim=1) == labels).float().mean() optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(vectorizer.parameters(), 1.0) optimizer.step() epoch_loss += loss.item() epoch_acc += acc.item() n_batches += 1 avg = lambda x: x / max(n_batches, 1) log["epoch"].append(epoch) log["loss"].append(avg(epoch_loss)) log["accuracy"].append(avg(epoch_acc)) if epoch % 10 == 0 or epoch == epochs - 1: print(f" Epoch {epoch:4d} | loss={avg(epoch_loss):.4f} acc={avg(epoch_acc):.3f}") torch.save(vectorizer.state_dict(), os.path.join(output_dir, "vectorizer.pt")) with open(os.path.join(output_dir, "phase2_log.json"), "w") as f: json.dump(log, f) print(f" Saved: {output_dir}/vectorizer.pt") return vectorizer, log def train_phase3a(data_dir, output_dir, vectorizer, epochs=50, lr=3e-4, batch_size=8): """Phase 3a: Train cross-encoder alignment (contrastive).""" print("\n=== Phase 3a: Cross-Encoder Alignment (Contrastive) ===") os.makedirs(output_dir, exist_ok=True) dataset = Phase3aDataset(data_dir, max_objects=6) loader = DataLoader(dataset, batch_size=batch_size, shuffle=True, drop_last=True) cross_enc = CrossEncoder(state_dim=256, embed_dim=256, action_dim=6) optimizer = torch.optim.Adam(cross_enc.parameters(), lr=lr) vectorizer.eval() log = {"epoch": [], "loss": [], "accuracy": []} for epoch in range(epochs): epoch_loss = 0 epoch_acc = 0 n_batches = 0 for batch in loader: B = batch["gt_poses"].shape[0] # Get scene vector from vectorizer phi = { 'z': torch.randn(B, 6, 64) * 0.1, 'poses': batch["gt_poses"], 'existence': batch["gt_existence"], } phi['z'][:, :, :4] = batch["gt_materials"][:, :6, :] graph = { 'edge_features': torch.randn(B, 6, 6, 32) * 0.01, 'contact_probs': batch["gt_contacts"], } with torch.no_grad(): s = vectorizer(phi, graph) scene_emb = cross_enc.encode_scene(s) text_emb = cross_enc.encode_text(batch["token_ids"], batch["attention_mask"]) loss = cross_enc.contrastive_loss(scene_emb, text_emb) # Accuracy sim = F.normalize(scene_emb, -1) @ F.normalize(text_emb, -1).T acc = (sim.argmax(1) == torch.arange(B)).float().mean() optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(cross_enc.parameters(), 1.0) optimizer.step() epoch_loss += loss.item() epoch_acc += acc.item() n_batches += 1 avg = lambda x: x / max(n_batches, 1) log["epoch"].append(epoch) log["loss"].append(avg(epoch_loss)) log["accuracy"].append(avg(epoch_acc)) if epoch % 10 == 0 or epoch == epochs - 1: print(f" Epoch {epoch:4d} | loss={avg(epoch_loss):.4f} acc={avg(epoch_acc):.3f}") torch.save(cross_enc.state_dict(), os.path.join(output_dir, "cross_encoder.pt")) with open(os.path.join(output_dir, "phase3a_log.json"), "w") as f: json.dump(log, f) print(f" Saved: {output_dir}/cross_encoder.pt") return cross_enc, log def train_phase3b(data_dir, output_dir, vectorizer, cross_enc, epochs=50, lr=1e-3, batch_size=8): """Phase 3b: Train action head (behavioral cloning).""" print("\n=== Phase 3b: Action Prediction (Behavioral Cloning) ===") os.makedirs(output_dir, exist_ok=True) dataset = Phase3bDataset(data_dir, max_objects=6) loader = DataLoader(dataset, batch_size=batch_size, shuffle=True, drop_last=True) # Freeze alignment layers, only train action head for p in cross_enc.scene_proj.parameters(): p.requires_grad = False for p in cross_enc.text_tf.parameters(): p.requires_grad = False for p in cross_enc.text_proj.parameters(): p.requires_grad = False optimizer = torch.optim.Adam( [p for p in cross_enc.parameters() if p.requires_grad], lr=lr) vectorizer.eval() log = {"epoch": [], "loss": [], "action_mae": []} for epoch in range(epochs): epoch_loss = 0 epoch_mae = 0 n_batches = 0 for batch in loader: B = batch["gt_poses"].shape[0] phi = { 'z': torch.randn(B, 6, 64) * 0.1, 'poses': batch["gt_poses"], 'existence': batch["gt_existence"], } phi['z'][:, :, :4] = batch["gt_materials"][:, :6, :] graph = { 'edge_features': torch.randn(B, 6, 6, 32) * 0.01, 'contact_probs': batch["gt_contacts"], } with torch.no_grad(): s = vectorizer(phi, graph) out = cross_enc(s, text_tokens=batch["token_ids"], text_mask=batch["attention_mask"]) loss = F.mse_loss(out["action"], batch["action"]) mae = (out["action"] - batch["action"]).abs().mean() optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(cross_enc.parameters(), 1.0) optimizer.step() epoch_loss += loss.item() epoch_mae += mae.item() n_batches += 1 avg = lambda x: x / max(n_batches, 1) log["epoch"].append(epoch) log["loss"].append(avg(epoch_loss)) log["action_mae"].append(avg(epoch_mae)) if epoch % 10 == 0 or epoch == epochs - 1: print(f" Epoch {epoch:4d} | loss={avg(epoch_loss):.4f} mae={avg(epoch_mae):.4f}") torch.save(cross_enc.state_dict(), os.path.join(output_dir, "cross_encoder_final.pt")) with open(os.path.join(output_dir, "phase3b_log.json"), "w") as f: json.dump(log, f) print(f" Saved: {output_dir}/cross_encoder_final.pt") return cross_enc, log # ============================================================ # MAIN # ============================================================ if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--phase", default="all", choices=["1","2","3a","3b","all"]) parser.add_argument("--data_dir", default="data") parser.add_argument("--output_dir", default="checkpoints") parser.add_argument("--epochs", type=int, default=50) parser.add_argument("--batch_size", type=int, default=4) parser.add_argument("--lr", type=float, default=3e-4) args = parser.parse_args() encoder = None vectorizer = None cross_enc = None if args.phase in ["1", "all"]: encoder, _ = train_phase1( f"{args.data_dir}/phase1", args.output_dir, epochs=args.epochs, lr=args.lr, batch_size=args.batch_size) if args.phase in ["2", "all"]: if encoder is None: encoder = EncodingModel(64, 6) ckpt = os.path.join(args.output_dir, "encoder.pt") if os.path.exists(ckpt): encoder.load_state_dict(torch.load(ckpt, weights_only=True)) vectorizer, _ = train_phase2( f"{args.data_dir}/phase2", args.output_dir, encoder, epochs=args.epochs, lr=args.lr, batch_size=args.batch_size) if args.phase in ["3a", "all"]: if vectorizer is None: vectorizer = StateVectorizer(64, 256) ckpt = os.path.join(args.output_dir, "vectorizer.pt") if os.path.exists(ckpt): vectorizer.load_state_dict(torch.load(ckpt, weights_only=True)) cross_enc, _ = train_phase3a( f"{args.data_dir}/phase3a", args.output_dir, vectorizer, epochs=args.epochs, lr=args.lr, batch_size=args.batch_size) if args.phase in ["3b", "all"]: if vectorizer is None: vectorizer = StateVectorizer(64, 256) ckpt = os.path.join(args.output_dir, "vectorizer.pt") if os.path.exists(ckpt): vectorizer.load_state_dict(torch.load(ckpt, weights_only=True)) if cross_enc is None: cross_enc = CrossEncoder(256, 256, 6) ckpt = os.path.join(args.output_dir, "cross_encoder.pt") if os.path.exists(ckpt): cross_enc.load_state_dict(torch.load(ckpt, weights_only=True)) train_phase3b( f"{args.data_dir}/phase3b", args.output_dir, vectorizer, cross_enc, epochs=args.epochs, lr=args.lr, batch_size=args.batch_size) print("\n=== Training Complete ===") print(f"Checkpoints: {args.output_dir}/")