--- 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])}") ```