Views
No views yet
| Parameter | Value |
|---|---|
| Base Model | cartesia-ai/Llamba-1B |
| Architecture | Mamba (SSM) |
| Layers | 16 |
| Hidden Dim | 2048 |
| Vocab Size | 128,256 |
| Training Data | WikiText-2 (1644 samples) |
| Training | 1 epoch |
| Metric | Value |
|---|---|
| Initial Loss | ~3.84 |
| Final Loss | ~1.43 |
| Layer 12 Loss | 0.775 |
| Layer 13 Loss | 0.714 |
| Layer 14 Loss | 0.443 |
| Layer 15 Loss | 0.203 |
1import torch
2import torch.nn as nn
3from huggingface_hub import hf_hub_download
4from transformers import AutoTokenizer
5from cartesia_pytorch.Llamba import LlambaLMHeadModel
6import json
7
8# Download files
9lens_path = hf_hub_download("Xeiroh/llamba-1b-tuned-lens", "lens.pt")
10config_path = hf_hub_download("Xeiroh/llamba-1b-tuned-lens", "config.json")
11
12# Load config
13with open(config_path) as f:
14 config = json.load(f)
15
16# Load base model
17model = LlambaLMHeadModel.from_pretrained(
18 "cartesia-ai/Llamba-1B",
19 torch_dtype=torch.bfloat16,
20 trust_remote_code=True,
21).cuda().eval()
22
23# Create lens module
24class MambaTunedLens(nn.Module):
25 def __init__(self, d_model, vocab_size, num_layers, bias=True,
26 unembed_weight=None, final_norm_weight=None, final_norm_eps=1e-5):
27 super().__init__()
28 self.d_model = d_model
29 self.vocab_size = vocab_size
30 self.num_layers = num_layers
31
32 self.translators = nn.ModuleList([
33 nn.Linear(d_model, d_model, bias=bias) for _ in range(num_layers)
34 ])
35
36 if unembed_weight is not None:
37 self.register_buffer("unembed", unembed_weight.clone())
38 if final_norm_weight is not None:
39 self.register_buffer("final_norm_weight", final_norm_weight.clone())
40 self.final_norm_eps = final_norm_eps
41
42 def forward_layer(self, hidden_state, layer_idx):
43 h = self.translators[layer_idx](hidden_state)
44 h = h * torch.rsqrt(h.pow(2).mean(-1, keepdim=True) + self.final_norm_eps)
45 h = h * self.final_norm_weight
46 return h @ self.unembed.T
47
48# Initialize lens with model weights
49unembed_weight = model.lm_head.weight.data
50final_norm_weight = model.backbone.final_layernorm.weight.data
51final_norm_eps = model.backbone.final_layernorm.variance_epsilon
52
53lens = MambaTunedLens(
54 d_model=config["d_model"],
55 vocab_size=config["vocab_size"],
56 num_layers=config["num_layers"],
57 bias=config["bias"],
58 unembed_weight=unembed_weight,
59 final_norm_weight=final_norm_weight,
60 final_norm_eps=final_norm_eps,
61).cuda()
62
63# Load trained weights
64lens.load_state_dict(torch.load(lens_path, weights_only=True))
65lens.eval()
66
67# Use the lens
68tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-3.1-8B-Instruct")
69inputs = tokenizer("The capital of France is", return_tensors="pt").to("cuda")
70
71with torch.no_grad():
72 outputs = model(inputs.input_ids, return_hidden_states=True)
73
74 # Get predictions from each layer
75 for layer_idx in range(config["num_layers"]):
76 hidden = outputs.all_hidden_states[layer_idx + 1]
77 logits = lens.forward_layer(hidden, layer_idx)
78 pred_token = logits[0, -1].argmax()
79 print(f"Layer {layer_idx}: {tokenizer.decode([pred_token])}")1@article{belrose2023eliciting,
2 title={Eliciting Latent Predictions from Transformers with the Tuned Lens},
3 author={Belrose, Nora and Furman, Zach and Smith, Logan and Halawi, Danny and Ostrovsky, Igor and McKinney, Lev and Biderman, Stella and Steinhardt, Jacob},
4 journal={arXiv preprint arXiv:2303.08112},
5 year={2023}
6}