File size: 9,094 Bytes
87e9895
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
"""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()