"""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()