Text Generation
Transformers
Safetensors
English
qwen3
long-context
sparse-attention
aha
l2a-style
reproducibility
conversational
text-generation-inference
Instructions to use keepsloading/icml_repro_scratch with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use keepsloading/icml_repro_scratch with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="keepsloading/icml_repro_scratch") messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# Load model directly from transformers import AutoTokenizer, AutoModelForCausalLM tokenizer = AutoTokenizer.from_pretrained("keepsloading/icml_repro_scratch") model = AutoModelForCausalLM.from_pretrained("keepsloading/icml_repro_scratch", device_map="auto") messages = [ {"role": "user", "content": "Who are you?"}, ] inputs = tokenizer.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", ).to(model.device) outputs = model.generate(**inputs, max_new_tokens=40) print(tokenizer.decode(outputs[0][inputs["input_ids"].shape[-1]:])) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use keepsloading/icml_repro_scratch with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "keepsloading/icml_repro_scratch" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "keepsloading/icml_repro_scratch", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/keepsloading/icml_repro_scratch
- SGLang
How to use keepsloading/icml_repro_scratch with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "keepsloading/icml_repro_scratch" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "keepsloading/icml_repro_scratch", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "keepsloading/icml_repro_scratch" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "keepsloading/icml_repro_scratch", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use keepsloading/icml_repro_scratch with Docker Model Runner:
docker model run hf.co/keepsloading/icml_repro_scratch
| from __future__ import annotations | |
| import json | |
| import math | |
| import os | |
| import tempfile | |
| import unittest | |
| from pathlib import Path | |
| from unittest import mock | |
| import torch | |
| from transformers import Qwen3Config, Qwen3ForCausalLM | |
| import modeling_aha_qwen3 as aha_module | |
| from modeling_aha_qwen3 import ( | |
| AHAQwen3Config, | |
| AHAQwen3ForCausalLM, | |
| aha_router_output_size, | |
| ) | |
| from router_training_utils import RowWiseAdamW, configure_gate_only | |
| def base_config() -> Qwen3Config: | |
| config = Qwen3Config( | |
| vocab_size=97, | |
| hidden_size=32, | |
| intermediate_size=64, | |
| num_hidden_layers=2, | |
| num_attention_heads=4, | |
| num_key_value_heads=2, | |
| head_dim=8, | |
| max_position_embeddings=64, | |
| attention_dropout=0.0, | |
| attention_bias=True, | |
| tie_word_embeddings=False, | |
| ) | |
| config._attn_implementation = "eager" | |
| return config | |
| def aha_config(granularity: str, *, force_gate_value=None) -> AHAQwen3Config: | |
| payload = base_config().to_dict() | |
| payload.update( | |
| aha_window_size=2, | |
| aha_local_kind="sliding_window", | |
| aha_mode="dynamic", | |
| aha_router_granularity=granularity, | |
| aha_force_gate_value=force_gate_value, | |
| aha_reg_weight=1.0, | |
| aha_ce_weight=0.0, | |
| model_type="aha_qwen3", | |
| ) | |
| config = AHAQwen3Config(**payload) | |
| config._attn_implementation = "eager" | |
| return config | |
| def copy_base_weights(base: Qwen3ForCausalLM, target: AHAQwen3ForCausalLM) -> None: | |
| source = base.state_dict() | |
| destination = target.state_dict() | |
| q_rows = target.config.num_attention_heads * target.config.head_dim | |
| with torch.no_grad(): | |
| for name, value in source.items(): | |
| if name not in destination: | |
| continue | |
| if destination[name].shape == value.shape: | |
| destination[name].copy_(value) | |
| elif name.endswith("self_attn.q_proj.weight"): | |
| destination[name][:q_rows].copy_(value) | |
| destination[name][q_rows:].zero_() | |
| elif name.endswith("self_attn.q_proj.bias"): | |
| destination[name][:q_rows].copy_(value) | |
| destination[name][q_rows:].zero_() | |
| else: | |
| raise AssertionError(f"unexpected shape mismatch for {name}") | |
| target.load_state_dict(destination) | |
| def aligned_models(): | |
| torch.manual_seed(7) | |
| base = Qwen3ForCausalLM(base_config()).eval() | |
| models = {} | |
| for granularity in ("token", "token_kv_head"): | |
| model = AHAQwen3ForCausalLM( | |
| aha_config(granularity, force_gate_value=1.0) | |
| ).eval() | |
| copy_base_weights(base, model) | |
| models[granularity] = model | |
| return base, models | |
| class RouterGranularityTest(unittest.TestCase): | |
| def test_auto_class_checkpoint_embeds_modeling_source(self): | |
| AHAQwen3Config.register_for_auto_class() | |
| AHAQwen3ForCausalLM.register_for_auto_class("AutoModelForCausalLM") | |
| with tempfile.TemporaryDirectory() as tmp: | |
| AHAQwen3ForCausalLM(aha_config("token")).save_pretrained( | |
| tmp, safe_serialization=True | |
| ) | |
| self.assertTrue((Path(tmp) / "modeling_aha_qwen3.py").exists()) | |
| def test_projection_and_effective_gate_shapes(self): | |
| _, models = aligned_models() | |
| input_ids = torch.tensor([[1, 2, 3, 4, 5]]) | |
| expected_q_rows = 4 * 8 | |
| for granularity, native_rows in (("token", 1), ("token_kv_head", 2)): | |
| with self.subTest(granularity=granularity): | |
| model = models[granularity] | |
| attention = model.model.layers[0].self_attn | |
| self.assertEqual(attention.q_proj.out_features, expected_q_rows + native_rows) | |
| self.assertEqual(attention.aha_router_outputs, native_rows) | |
| output = model.model(input_ids=input_ids, use_cache=False) | |
| self.assertEqual(len(output.all_gate_soft), 2) | |
| self.assertEqual(output.all_gate_soft[0].shape, (1, 5, 2)) | |
| self.assertEqual(output.all_gate_hard[0].shape, (1, 5, 2)) | |
| def test_force_open_matches_full_attention_logits(self): | |
| base, models = aligned_models() | |
| input_ids = torch.tensor([[1, 4, 2, 8, 3, 9]]) | |
| with torch.no_grad(): | |
| expected = base(input_ids=input_ids, use_cache=False).logits | |
| for granularity, model in models.items(): | |
| with self.subTest(granularity=granularity): | |
| actual = model(input_ids=input_ids, use_cache=False).logits | |
| torch.testing.assert_close(actual, expected, atol=1e-6, rtol=1e-5) | |
| def test_force_closed_uses_identical_local_branch(self): | |
| _, models = aligned_models() | |
| input_ids = torch.tensor([[1, 4, 2, 8, 3, 9]]) | |
| for model in models.values(): | |
| model.config.aha_force_gate_value = 0.0 | |
| with torch.no_grad(): | |
| token_logits = models["token"](input_ids=input_ids, use_cache=False).logits | |
| head_logits = models["token_kv_head"]( | |
| input_ids=input_ids, use_cache=False | |
| ).logits | |
| torch.testing.assert_close(token_logits, head_logits, atol=1e-6, rtol=1e-5) | |
| def test_regularizer_uses_effective_kv_head_denominator(self): | |
| _, models = aligned_models() | |
| probability = 0.73 | |
| bias = math.log(probability / (1.0 - probability)) | |
| input_ids = torch.tensor([[1, 4, 2, 8, 3, 9]]) | |
| outputs = {} | |
| for granularity, model in models.items(): | |
| model.train() | |
| model.config.aha_force_gate_value = None | |
| q_rows = model.config.num_attention_heads * model.config.head_dim | |
| with torch.no_grad(): | |
| for layer in model.model.layers: | |
| layer.self_attn.q_proj.weight[q_rows:].zero_() | |
| layer.self_attn.q_proj.bias[q_rows:].fill_(bias) | |
| outputs[granularity] = model( | |
| input_ids=input_ids, labels=input_ids.clone(), use_cache=False | |
| ) | |
| torch.testing.assert_close( | |
| outputs["token"].gate_soft_mean, | |
| outputs["token_kv_head"].gate_soft_mean, | |
| ) | |
| torch.testing.assert_close( | |
| outputs["token"].gate_aux_loss, | |
| outputs["token_kv_head"].gate_aux_loss, | |
| ) | |
| self.assertAlmostEqual(outputs["token"].gate_soft_mean.item(), probability, places=6) | |
| def test_gate_only_updates_only_appended_rows(self): | |
| for granularity, expected_rows in (("token", 1), ("token_kv_head", 2)): | |
| with self.subTest(granularity=granularity): | |
| model = AHAQwen3ForCausalLM(aha_config(granularity)) | |
| setup = configure_gate_only(model) | |
| self.assertEqual(setup.gate_rows, expected_rows) | |
| q_proj = model.model.layers[0].self_attn.q_proj | |
| before = q_proj.weight.detach().clone() | |
| q_proj.weight.sum().backward() | |
| self.assertEqual( | |
| torch.count_nonzero(q_proj.weight.grad[: setup.q_rows]).item(), 0 | |
| ) | |
| self.assertGreater( | |
| torch.count_nonzero(q_proj.weight.grad[setup.q_rows :]).item(), 0 | |
| ) | |
| torch.optim.SGD(setup.parameters, lr=0.1).step() | |
| torch.testing.assert_close( | |
| q_proj.weight[: setup.q_rows], before[: setup.q_rows] | |
| ) | |
| self.assertFalse( | |
| torch.equal(q_proj.weight[setup.q_rows :], before[setup.q_rows :]) | |
| ) | |
| def test_rowwise_adamw_applies_real_ten_x_lr_ratio(self): | |
| parameter = torch.nn.Parameter(torch.zeros(3, 1)) | |
| optimizer = RowWiseAdamW( | |
| [{"params": [parameter], "lr": 0.1}], | |
| row_scales=[(parameter, 2, 0.1)], | |
| weight_decay=0.0, | |
| betas=(0.9, 0.999), | |
| ) | |
| parameter.grad = torch.ones_like(parameter) | |
| optimizer.step() | |
| backbone_update = parameter[:2].abs().mean().item() | |
| gate_update = parameter[2:].abs().mean().item() | |
| self.assertAlmostEqual(gate_update / backbone_update, 10.0, places=5) | |
| def test_one_step_smoke_is_finite_and_reloadable(self): | |
| input_ids = torch.tensor([[1, 4, 2, 8, 3, 9]]) | |
| for granularity in ("token", "token_kv_head"): | |
| with self.subTest(granularity=granularity), tempfile.TemporaryDirectory() as tmp: | |
| model = AHAQwen3ForCausalLM(aha_config(granularity)).train() | |
| setup = configure_gate_only(model) | |
| optimizer = torch.optim.AdamW(setup.parameters, lr=3e-5) | |
| output = model( | |
| input_ids=input_ids, labels=input_ids.clone(), use_cache=False | |
| ) | |
| self.assertTrue(torch.isfinite(output.loss).item()) | |
| output.loss.backward() | |
| self.assertTrue( | |
| all( | |
| parameter.grad is None | |
| or torch.isfinite(parameter.grad).all().item() | |
| for parameter in setup.parameters | |
| ) | |
| ) | |
| optimizer.step() | |
| model.save_pretrained(tmp, safe_serialization=True) | |
| reloaded = AHAQwen3ForCausalLM.from_pretrained_aha( | |
| tmp, torch_dtype=torch.float32, attn_implementation="eager" | |
| ) | |
| self.assertEqual( | |
| reloaded.config.aha_router_granularity, granularity | |
| ) | |
| def test_legacy_default_and_save_load_roundtrip(self): | |
| legacy = AHAQwen3Config(**base_config().to_dict()) | |
| self.assertEqual(legacy.aha_router_granularity, "token_kv_head") | |
| self.assertEqual(aha_router_output_size(legacy), legacy.num_key_value_heads) | |
| with tempfile.TemporaryDirectory() as tmp: | |
| model = AHAQwen3ForCausalLM(aha_config("token")) | |
| model.save_pretrained(tmp, safe_serialization=True) | |
| loaded = AHAQwen3ForCausalLM.from_pretrained_aha( | |
| tmp, torch_dtype=torch.float32, attn_implementation="eager" | |
| ) | |
| self.assertEqual(loaded.config.aha_router_granularity, "token") | |
| self.assertEqual(loaded.model.layers[0].self_attn.aha_router_outputs, 1) | |
| def test_sparsity_tracker_records_native_and_effective_counts(self): | |
| with tempfile.TemporaryDirectory() as tmp: | |
| output = os.path.join(tmp, "sparsity.json") | |
| with mock.patch.dict(os.environ, {"AHA_SPARSITY_STATS_PATH": output}): | |
| tracker = aha_module._AHAInferenceSparsityTracker() | |
| gate_hard = torch.tensor([[[1.0, 1.0], [0.0, 0.0]]]) | |
| gate_soft = torch.tensor([[[0.8, 0.8], [0.2, 0.2]]]) | |
| tracker.update( | |
| gate_hard, | |
| gate_soft, | |
| layer_idx=0, | |
| phase="prefill", | |
| router_granularity="token", | |
| native_router_width=1, | |
| ) | |
| tracker.update( | |
| gate_hard[:, :1], | |
| gate_soft[:, :1], | |
| layer_idx=0, | |
| phase="decode", | |
| router_granularity="token", | |
| native_router_width=1, | |
| ) | |
| tracker.write_stats() | |
| payload = json.loads(Path(output).read_text(encoding="utf-8")) | |
| self.assertEqual(payload["router_granularity"], "token") | |
| self.assertEqual(payload["native_router_decisions"], 3) | |
| self.assertEqual(payload["effective_router_decisions"], 6) | |
| self.assertAlmostEqual(payload["sparsity"], 1 / 3) | |
| self.assertEqual(payload["by_phase"]["decode"]["total_decisions"], 2) | |
| if __name__ == "__main__": | |
| unittest.main() | |