| --- |
| license: cc-by-nc-4.0 |
| datasets: |
| - Salesforce/wikitext |
| base_model: |
| - CohereLabs/c4ai-command-r-v01 |
| tags: |
| - jlens, |
| - jacobian-lens, |
| - interpretability |
| --- |
| |
| ```bash |
| git clone https://github.com/anthropics/jacobian-lens |
| cd jacobian-lens |
| ``` |
|
|
| ## Load in BF16 (Requires 96GB VRAM) |
|
|
| nfp4 works with bitsandbytes |
|
|
| ```python |
| import jlens |
| import torch |
| import transformers |
| jlens.configure_logging() |
| |
| # Config |
| LENS_REPO="gghfez/c4ai-command-r-v01-jacobian-lens" |
| MODEL_NAME = "CohereLabs/c4ai-command-r-v01" |
| LENS_FILE="c4ai-command-r-v01_jlens.pt" |
| |
| # can use flash-attn2 if instead of spda |
| hf_model = transformers.AutoModelForCausalLM.from_pretrained( |
| MODEL_NAME, dtype=torch.bfloat16, attn_implementation="sdpa" |
| ).cuda() |
| tokenizer = transformers.AutoTokenizer.from_pretrained(MODEL_NAME) |
| |
| |
| hf_model.gradient_checkpointing_enable() |
| hf_model.config.use_cache = False |
| hf_model.requires_grad_(False) |
| |
| lens = jlens.JacobianLens.from_pretrained(LENS_REPO, filename=LENS_FILE) |
| lens |
| #JacobianLens(d_model=8192, n_prompts=100, source_layers=[0..38] (39 layers)) |
| |
| model = jlens.from_hf(hf_model, tokenizer) |
| model |
| #HFLensModel(CohereForCausalLM, n_layers=40, d_model=8192) |
| |
| |
| ``` |
|
|
|
|
|
|
| ## Example logits vs jlens |
|
|
| ⚠️ J-lens shows "degenerate" (not just "degenerative") output for this model |
|
|
| ```python |
| prompt = """Hey Gemma, what do you want most in the world.""" |
| |
| layers={k: None for k in range(3, 39)} |
| |
| logit_lens, _, _ = lens.apply(model, prompt, layers=layers, positions=[-2], use_jacobian=False) |
| |
| jlens_logits, model_logits, _ = lens.apply(model, prompt, layers=layers, positions=[-2]) |
| |
| |
| def top10(logits): |
| return [tokenizer.decode([t]) for t in logits.topk(10).indices] |
| |
| def top5(logits): |
| return [tokenizer.decode([t]) for t in logits.topk(5).indices] |
| |
| print("-"*6, "Jlens", "-"*6) |
| for layer in layers: |
| print(f"L{layer:>3} J-lens: {top5(jlens_logits[layer][0])}") |
| |
| print("-"*6, "Logits", "-"*6) |
| for layer in layers: |
| print(f"L{layer:>3} logit-lens: {top5(logit_lens[layer][0])}") |
| |
| print("-"*6) |
| print(f"Model prediction: {top5(model_logits[0])}") |
| ``` |