Nirav Madhani
Review updated training assets
545e799
Raw
History Blame Contribute Delete
30 kB
"""
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}/")