File size: 1,125 Bytes
3b22208
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Minimal, self-contained loader for the V-JEPA 2 archetypal sparse transcoder.

Requires: torch, safetensors, and modeling_archetypal_transcoder.py (this repo).
"""
import torch
from safetensors.torch import load_file
from modeling_archetypal_transcoder import ArchetypalTranscoder, ArchetypalConfig

def load_transcoder(weights="model.safetensors", device="cpu"):
    sd = load_file(weights)
    cfg = ArchetypalConfig(
        d_in=768, d_hidden=12288, n_pool=32000, top_k=512,
        init_eye=True, delta=0.5, encoder_topk=None,
    )
    # data_pool (K-means centroids) is bundled in the weights; pass it to the ctor.
    model = ArchetypalTranscoder(cfg, data_pool=sd["data_pool"]).to(device)
    model.load_state_dict(sd)
    model.eval()
    return model

if __name__ == "__main__":
    m = load_transcoder()
    # x = V-JEPA 2.1 block-10 FFN *input* activations, shape [N, 768], fp32.
    x = torch.randn(4, 768)
    x_hat, z = m(x)          # x_hat: predicted FFN output [N,768]; z: sparse codes [N,12288]
    print("x_hat", tuple(x_hat.shape), "| z", tuple(z.shape), "| L0", (z>0).float().sum(-1).mean().item())