Instructions to use jamie33/mind3d-trellis2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Trellis
How to use jamie33/mind3d-trellis2 with Trellis:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
Download code/mapper/train_mapper.py from jamie33/mind3d-trellis2: direct link, hf CLI and curl.
- Browser
- Download file 9.09 kB
-
https://huggingface.co/jamie33/mind3d-trellis2/resolve/main/code/mapper/train_mapper.py
- Command line
-
hf download hf://jamie33/mind3d-trellis2/code/mapper/train_mapper.py
-
curl -L -o train_mapper.py https://huggingface.co/jamie33/mind3d-trellis2/resolve/main/code/mapper/train_mapper.py
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) | |
| 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() | |