mind3d-trellis2 / code /mapper /train_mapper.py
jamie33's picture
MinD-3D + TRELLIS.2 sub-01 experiments: code, metrics, logs, report
87e9895 verified
Raw History Blame Contribute Delete
9.09 kB
"""fMRI (frozen MinD-3D encoder features) -> TRELLIS.2 DINOv3 conditioning tokens."""
import os
import json
import math
import time
import argparse
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
FEAT_DIR = "/home/hubin/data/fMRI-Shape/mapper_feats"
TOKEN_DIR = "/home/hubin/data/fMRI-Shape/trellis2_cond/tokens"
TOKEN_FRAMES = [0, 24, 48, 72, 96, 120, 144, 168]
N_TOKENS, TOKEN_DIM = 1029, 1024
N_FMRI_FRAMES, N_USED_FRAMES, FMRI_TOKENS = 10, 6, 257
TEST_FRAMES = [2, 3, 4, 5, 6, 7]
class Mapper(nn.Module):
def __init__(self, dim, depth, heads, dropout, mean_token):
super().__init__()
self.in_proj = nn.Sequential(nn.LayerNorm(1024), nn.Linear(1024, dim))
self.frame_emb = nn.Embedding(N_FMRI_FRAMES, dim)
self.pos_emb = nn.Parameter(torch.zeros(FMRI_TOKENS, dim))
self.queries = nn.Parameter(torch.randn(N_TOKENS, dim) * 0.02)
layer = nn.TransformerDecoderLayer(dim, heads, 4 * dim, dropout, activation="gelu",
batch_first=True, norm_first=True)
self.decoder = nn.TransformerDecoder(layer, depth)
self.out = nn.Sequential(nn.LayerNorm(dim), nn.Linear(dim, TOKEN_DIM))
nn.init.zeros_(self.out[1].weight)
nn.init.zeros_(self.out[1].bias)
self.register_buffer("mean_token", mean_token)
def forward(self, feats, frame_idx, mem_drop=0.0):
# feats: [B, F, 257, 1024], frame_idx: [B, F]
b, f = frame_idx.shape
mem = self.in_proj(feats) + self.pos_emb + self.frame_emb(frame_idx)[:, :, None]
mem = mem.flatten(1, 2)
if self.training and mem_drop > 0:
keep = int(mem.shape[1] * (1 - mem_drop))
idx = torch.rand(b, mem.shape[1], device=mem.device).argsort(1)[:, :keep]
mem = torch.gather(mem, 1, idx[..., None].expand(-1, -1, mem.shape[-1]))
x = self.decoder(self.queries.expand(b, -1, -1), mem)
return F.layer_norm(self.out(x) + self.mean_token, (TOKEN_DIM,))
def load_tokens(ids, frame):
k = TOKEN_FRAMES.index(frame)
return np.stack([np.load(f"{TOKEN_DIR}/{i}.npy", mmap_mode="r")[k] for i in ids])
def pooled(x):
return F.normalize(x.float().mean(1), dim=-1)
@torch.no_grad()
def evaluate(model, feats, frames, targets, bs=16):
model.eval()
preds = []
for s in range(0, len(feats), bs):
fi = torch.tensor(frames, device="cuda").expand(min(bs, len(feats) - s), -1)
with torch.autocast("cuda", dtype=torch.bfloat16):
preds.append(model(feats[s:s + bs][:, frames].float(), fi).float())
model.train()
preds = torch.cat(preds)
return preds, token_metrics(preds, targets)
def token_metrics(preds, targets):
targets = targets.float()
cos = F.cosine_similarity(preds, targets, dim=-1).mean().item()
mse = F.mse_loss(preds, targets).item()
sim = pooled(preds) @ pooled(targets).T
rank = (sim > sim.diag()[:, None]).sum(1)
return {"cos": cos, "mse": mse, "top1": (rank < 1).float().mean().item(),
"top5": (rank < 5).float().mean().item(), "n": len(preds)}
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--out_dir", default="/home/hubin/trellis_work/outputs/mapper_v0")
parser.add_argument("--sub_id", default="0001")
parser.add_argument("--target_frame", type=int, default=24)
parser.add_argument("--dim", type=int, default=768)
parser.add_argument("--depth", type=int, default=6)
parser.add_argument("--heads", type=int, default=12)
parser.add_argument("--dropout", type=float, default=0.1)
parser.add_argument("--mem_drop", type=float, default=0.2)
parser.add_argument("--feat_noise", type=float, default=0.1)
parser.add_argument("--contrastive", type=float, default=0.1)
parser.add_argument("--tau", type=float, default=0.05)
parser.add_argument("--lr", type=float, default=3e-4)
parser.add_argument("--wd", type=float, default=0.05)
parser.add_argument("--bs", type=int, default=32)
parser.add_argument("--steps", type=int, default=6000)
parser.add_argument("--n_val", type=int, default=100)
parser.add_argument("--eval_every", type=int, default=500)
parser.add_argument("--seed", type=int, default=0)
args = parser.parse_args()
os.makedirs(args.out_dir, exist_ok=True)
torch.manual_seed(args.seed)
rng = np.random.default_rng(args.seed)
train_ids = open(f"{FEAT_DIR}/train_ids.txt").read().split()
test_ids = open(f"{FEAT_DIR}/test_ids.txt").read().split()
train_feats = torch.from_numpy(np.load(f"{FEAT_DIR}/sub{args.sub_id}_train_feats.npy")).cuda()
test_feats = torch.from_numpy(np.load(f"{FEAT_DIR}/sub{args.sub_id}_test_feats.npy")).cuda()
train_tok = torch.from_numpy(load_tokens(train_ids, args.target_frame)).cuda()
test_tok = torch.from_numpy(load_tokens(test_ids, args.target_frame)).cuda()
perm = rng.permutation(len(train_ids))
val_idx = torch.from_numpy(np.sort(perm[:args.n_val])).cuda()
fit_idx = torch.from_numpy(np.sort(perm[args.n_val:])).cuda()
fit_cats = np.array([train_ids[i].split("/")[0] for i in fit_idx.tolist()])
mean_token = train_tok[fit_idx].float().mean(0)
ref = {}
for name, idx in [("val", val_idx), ("test", None)]:
tgt = train_tok[idx] if idx is not None else test_tok
ids = [train_ids[i] for i in idx.tolist()] if idx is not None else test_ids
mean_pred = F.layer_norm(mean_token, (TOKEN_DIM,)).expand(len(tgt), -1, -1)
cat_means = {c: train_tok[fit_idx[torch.from_numpy(np.where(fit_cats == c)[0]).cuda()]].float().mean(0)
for c in set(fit_cats)}
cat_pred = torch.stack([F.layer_norm(cat_means[i.split("/")[0]], (TOKEN_DIM,)) for i in ids])
ref[name] = {"train_mean": token_metrics(mean_pred, tgt),
"category_mean_oracle": token_metrics(cat_pred, tgt)}
print(json.dumps(ref, indent=1), flush=True)
model = Mapper(args.dim, args.depth, args.heads, args.dropout, mean_token).cuda()
print(f"params {sum(p.numel() for p in model.parameters()) / 1e6:.1f}M", flush=True)
opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.wd, betas=(0.9, 0.98))
warmup = 300
sched = torch.optim.lr_scheduler.LambdaLR(
opt, lambda s: min(1, (s + 1) / warmup) * 0.5 * (1 + math.cos(math.pi * min(s, args.steps) / args.steps)))
log, best = [], (-1, None)
t0 = time.time()
for step in range(1, args.steps + 1):
bidx = fit_idx[torch.randint(len(fit_idx), (args.bs,), device="cuda")]
frames = torch.rand(args.bs, N_FMRI_FRAMES, device="cuda").argsort(1)[:, :N_USED_FRAMES].sort(1).values
feats = torch.gather(train_feats[bidx], 1, frames[:, :, None, None].expand(-1, -1, FMRI_TOKENS, 1024)).float()
feats = feats + args.feat_noise * torch.randn_like(feats)
tgt = train_tok[bidx].float()
with torch.autocast("cuda", dtype=torch.bfloat16):
pred = model(feats, frames, args.mem_drop)
pred = pred.float()
loss_tok = F.mse_loss(pred, tgt)
logits = pooled(pred) @ pooled(tgt).T / args.tau
labels = torch.arange(args.bs, device="cuda")
loss_con = (F.cross_entropy(logits, labels) + F.cross_entropy(logits.T, labels)) / 2
loss = loss_tok + args.contrastive * loss_con
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
sched.step()
if step % 100 == 0:
print(f"step {step} loss_tok {loss_tok.item():.4f} loss_con {loss_con.item():.3f} "
f"lr {sched.get_last_lr()[0]:.2e} {time.time() - t0:.0f}s", flush=True)
if step % args.eval_every == 0 or step == args.steps:
_, val_m = evaluate(model, train_feats[val_idx], TEST_FRAMES, train_tok[val_idx])
log.append({"step": step, "val": val_m})
print(f" [val] {val_m}", flush=True)
if val_m["cos"] > best[0]:
best = (val_m["cos"], step)
torch.save(model.state_dict(), f"{args.out_dir}/best.pt")
model.load_state_dict(torch.load(f"{args.out_dir}/best.pt"))
test_pred, test_m = evaluate(model, test_feats, TEST_FRAMES, test_tok)
train_pool = pooled(train_tok[fit_idx])
nn_idx = (pooled(test_pred) @ train_pool.T).argmax(1)
nn_ids = [train_ids[fit_idx[i]] for i in nn_idx.tolist()]
nn_cat_acc = float(np.mean([a.split("/")[0] == b.split("/")[0] for a, b in zip(nn_ids, test_ids)]))
np.save(f"{args.out_dir}/test_pred_tokens.npy", test_pred.half().cpu().numpy())
result = {"args": vars(args), "best_step": best[1], "reference": ref, "test": test_m,
"test_nn_category_acc": nn_cat_acc, "test_nn_ids": dict(zip(test_ids, nn_ids)), "log": log}
with open(f"{args.out_dir}/result.json", "w") as f:
json.dump(result, f, indent=1)
print("TEST", test_m, "nn_category_acc", nn_cat_acc, "best_step", best[1], flush=True)
if __name__ == "__main__":
main()