| """ |
| 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 |
|
|
| |
| |
|
|
|
|
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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"]) |
| |
| |
| 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 = 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)), |
| } |
| |
| |
| 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) |
| |
| |
| 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): |
| |
| return sum(e["n_steps"] for e in self.manifest) |
| |
| def __getitem__(self, idx): |
| |
| 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"]), |
| } |
|
|
|
|
| |
| |
| |
|
|
| 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 = F.mse_loss(out['phi']['poses'][:, :, :3], batch["gt_poses"][:, :, :3]) |
| |
| |
| exist_loss = F.binary_cross_entropy(out['phi']['existence'], batch["gt_existence"]) |
| |
| |
| 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}") |
| |
| |
| 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] |
| |
| |
| |
| phi_pos = { |
| 'z': torch.randn(B, 6, 64) * 0.1, |
| '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"], |
| } |
| |
| |
| 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 |
| |
| |
| phi_neg = { |
| 'z': torch.randn(B, 6, 64) * 0.1, |
| 'poses': batch["negative"]["gt_poses"], |
| 'existence': batch["negative"]["gt_existence"], |
| } |
| |
| |
| mat_pos = batch["positive"]["gt_materials"] |
| mat_neg = batch["negative"]["gt_materials"] |
| |
| phi_pos['z'][:, :, :4] = mat_pos[:, :6, :] |
| 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) |
| s_b = vectorizer(phi_pos2, graph_pos2) |
| |
| |
| 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 |
| labels = torch.arange(B) |
| loss = (F.cross_entropy(logits, labels) + F.cross_entropy(logits.T, labels)) / 2 |
| |
| |
| 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] |
| |
| |
| 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) |
| |
| |
| 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) |
| |
| |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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}/") |
|
|