File size: 4,668 Bytes
1b929be
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Distil BAAI/bge-m3 dense embeddings into a smaller encoder (same 1024-d vector space).
python emb/train_emb.py --base jhu-clsp/mmBERT-small --teacher DIR[,DIR] --out OUT
Student = encoder + mean pooling + linear projection (hidden -> 1024, saved as OUT/projection.pt), L2-normalised.
Loss = (1 - cos(student, teacher)) + mse_w * ||student - teacher||^2 + sim_w * in-batch similarity-matrix MSE (keeps the
relative geometry). Val = 0.5% of texts: mean cosine to the teacher and top-1 in-batch retrieval agreement.
"""
import argparse, json, os, random, time
import numpy as np, pandas as pd, torch
import torch.nn.functional as F
from transformers import AutoModel, AutoTokenizer, get_cosine_schedule_with_warmup

ap = argparse.ArgumentParser()
ap.add_argument("--base", default="jhu-clsp/mmBERT-small"); ap.add_argument("--teacher", required=True); ap.add_argument("--out", required=True)
ap.add_argument("--max_len", type=int, default=256); ap.add_argument("--bs", type=int, default=256); ap.add_argument("--lr", type=float, default=1e-4)
ap.add_argument("--epochs", type=float, default=2.0); ap.add_argument("--warmup", type=float, default=0.03)
ap.add_argument("--mse_w", type=float, default=1.0); ap.add_argument("--sim_w", type=float, default=1.0)
ap.add_argument("--seed", type=int, default=0); ap.add_argument("--eval_every", type=int, default=2000)
args = ap.parse_args()
random.seed(args.seed); np.random.seed(args.seed); torch.manual_seed(args.seed)
os.makedirs(args.out, exist_ok=True)
texts, embs = [], []
for d in args.teacher.split(","):
    texts += pd.read_parquet(f"{d}/texts.parquet").text.tolist(); embs.append(np.load(f"{d}/emb.npy", mmap_mode="r"))
E = np.concatenate([np.asarray(e) for e in embs]).astype(np.float16)
n = len(texts); rng = np.random.RandomState(0); isval = rng.rand(n) < 0.005
tr_idx, va_idx = np.where(~isval)[0], np.where(isval)[0][:4096]
print("texts", n, "train", len(tr_idx), "val", len(va_idx), flush=True)
tok = AutoTokenizer.from_pretrained(args.base); enc = AutoModel.from_pretrained(args.base).cuda()
proj = torch.nn.Linear(enc.config.hidden_size, 1024, bias=False).cuda()
params = list(enc.parameters()) + list(proj.parameters())


def embed(idx):
    e = tok([texts[i] for i in idx], truncation=True, max_length=args.max_len, padding=True, return_tensors="pt").to("cuda")
    with torch.autocast("cuda", dtype=torch.bfloat16):
        h = enc(input_ids=e["input_ids"], attention_mask=e["attention_mask"]).last_hidden_state
    a = e["attention_mask"].unsqueeze(-1).float(); v = (h.float() * a).sum(1) / a.sum(1)
    return F.normalize(proj(v), dim=-1)


def loss_fn(s, t):
    cos = (s * t).sum(-1)
    return (1 - cos).mean() + args.mse_w * ((s - t) ** 2).sum(-1).mean() + args.sim_w * ((s @ s.T - t @ t.T) ** 2).mean() * 100


@torch.no_grad()
def evaluate():
    enc.eval(); S = []
    for s in range(0, len(va_idx), 256):
        S.append(embed(va_idx[s:s + 256]))
    enc.train(); S = torch.cat(S); T = torch.tensor(E[va_idx], dtype=torch.float32, device="cuda")
    top = ((S @ T.T).argmax(1) == torch.arange(len(S), device="cuda")).float().mean().item()
    return dict(cos=float((S * T).sum(-1).mean().item()), top1=top)


steps = int(len(tr_idx) / args.bs * args.epochs)
opt = torch.optim.AdamW(params, lr=args.lr, weight_decay=0.01, betas=(0.9, 0.98), eps=1e-6, fused=True)
sch = get_cosine_schedule_with_warmup(opt, int(args.warmup * steps), steps)
print("steps", steps, flush=True)
step, best, log, t0 = 0, -1, [], time.time()
while step < steps:
    perm = np.random.permutation(tr_idx)
    for s in range(0, len(perm) - args.bs + 1, args.bs):
        b = perm[s:s + args.bs]
        st = embed(b); tt = torch.tensor(E[b], dtype=torch.float32, device="cuda")
        loss = loss_fn(st, tt); loss.backward(); torch.nn.utils.clip_grad_norm_(params, 1.0)
        opt.step(); sch.step(); opt.zero_grad(set_to_none=True); step += 1
        if step % 200 == 0:
            print(f"step {step}/{steps} loss {loss.item():.4f} {time.time()-t0:.0f}s", flush=True)
        if step % args.eval_every == 0 or step == steps:
            r = evaluate(); log.append(dict(step=step, **r)); print("EVAL", step, json.dumps(r), flush=True)
            if r["cos"] > best:
                best = r["cos"]; enc.save_pretrained(args.out); tok.save_pretrained(args.out); torch.save(proj.weight.detach().T.contiguous().cpu(), f"{args.out}/projection.pt")
                json.dump(dict(step=step, **r), open(f"{args.out}/val_metrics.json", "w"), indent=1)
        if step >= steps:
            break
json.dump(dict(args=vars(args), log=log), open(f"{args.out}/train_log.json", "w"), indent=1)
print("done; best val cos", best)