interp-engine 1.7.3__tar.gz → 1.9.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (180) hide show
  1. {interp_engine-1.7.3 → interp_engine-1.9.0}/PKG-INFO +1 -1
  2. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/probe.py +8 -2
  3. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/run_all.sh +29 -5
  4. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/run_bench.py +25 -6
  5. {interp_engine-1.7.3 → interp_engine-1.9.0}/docs/AGENT_INTEGRATION.md +7 -5
  6. {interp_engine-1.7.3 → interp_engine-1.9.0}/docs/INTERNALS.md +9 -1
  7. {interp_engine-1.7.3 → interp_engine-1.9.0}/docs/SUPPORTED_POINTS.md +14 -11
  8. {interp_engine-1.7.3 → interp_engine-1.9.0}/docs/USAGE.md +10 -0
  9. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/capture.py +2 -2
  10. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/facts.py +19 -0
  11. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/load.py +73 -5
  12. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/memory.py +318 -31
  13. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/model.py +153 -12
  14. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/points.py +2 -1
  15. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/protocol.py +1 -1
  16. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_backend.py +34 -15
  17. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/__init__.py +2 -0
  18. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/_hooks.py +15 -3
  19. interp_engine-1.9.0/interp_engine/vllm_capture/_tp.py +147 -0
  20. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/_tree.py +104 -20
  21. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/attn.py +43 -5
  22. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/capture.py +16 -2
  23. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/requests.py +36 -5
  24. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/static.py +70 -7
  25. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_plugin.py +10 -0
  26. {interp_engine-1.7.3 → interp_engine-1.9.0}/pyproject.toml +2 -1
  27. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_facts.py +7 -0
  28. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_gpu_sizer.py +41 -6
  29. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_load.py +179 -1
  30. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_memory.py +234 -1
  31. interp_engine-1.9.0/tests/test_multigpu.py +236 -0
  32. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_multimodal_arch.py +2 -1
  33. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_static_set.py +67 -0
  34. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_kv_isolation.py +2 -0
  35. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_new_points.py +112 -21
  36. {interp_engine-1.7.3 → interp_engine-1.9.0}/.gitignore +0 -0
  37. {interp_engine-1.7.3 → interp_engine-1.9.0}/LICENSE +0 -0
  38. {interp_engine-1.7.3 → interp_engine-1.9.0}/README.md +0 -0
  39. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/README.md +0 -0
  40. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/__init__.py +0 -0
  41. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/bench_spec.py +0 -0
  42. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/cells.py +0 -0
  43. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/probe_lens_stream.py +0 -0
  44. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/publish.py +0 -0
  45. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/report_bench.py +0 -0
  46. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/deepseek-v4-flash-0731__eager.json +0 -0
  47. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/deepseek-v4-flash-0731__vllm-cudagraph.json +0 -0
  48. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/deepseek-v4-flash-0731__vllm-static.json +0 -0
  49. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/deepseek-v4-flash-0731__vllm.json +0 -0
  50. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/gemma-2-2b__eager.json +0 -0
  51. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/gemma-2-2b__vllm-cudagraph.json +0 -0
  52. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/gemma-2-2b__vllm-static.json +0 -0
  53. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/gemma-2-2b__vllm.json +0 -0
  54. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/llama-3.1-8b__eager.json +0 -0
  55. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/llama-3.1-8b__vllm-cudagraph.json +0 -0
  56. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/llama-3.1-8b__vllm-static.json +0 -0
  57. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/llama-3.1-8b__vllm.json +0 -0
  58. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/qwen3-4b__eager.json +0 -0
  59. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/qwen3-4b__vllm-cudagraph.json +0 -0
  60. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/qwen3-4b__vllm-static.json +0 -0
  61. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/qwen3-4b__vllm.json +0 -0
  62. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/qwen3.8-27b__eager.json +0 -0
  63. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/qwen3.8-27b__vllm-cudagraph.json +0 -0
  64. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/qwen3.8-27b__vllm-static.json +0 -0
  65. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results/qwen3.8-27b__vllm.json +0 -0
  66. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/results-latest.md +0 -0
  67. {interp_engine-1.7.3 → interp_engine-1.9.0}/benchmarks/workloads.py +0 -0
  68. {interp_engine-1.7.3 → interp_engine-1.9.0}/docs/ARCHITECTURE_QUIRKS.md +0 -0
  69. {interp_engine-1.7.3 → interp_engine-1.9.0}/docs/COMPATIBILITY.md +0 -0
  70. {interp_engine-1.7.3 → interp_engine-1.9.0}/docs/ENGINE_HOOK_MAPPINGS.md +0 -0
  71. {interp_engine-1.7.3 → interp_engine-1.9.0}/docs/GRADIENTS.md +0 -0
  72. {interp_engine-1.7.3 → interp_engine-1.9.0}/docs/PERFORMANCE.md +0 -0
  73. {interp_engine-1.7.3 → interp_engine-1.9.0}/docs/PORTING.md +0 -0
  74. {interp_engine-1.7.3 → interp_engine-1.9.0}/docs/README.md +0 -0
  75. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/__init__.py +0 -0
  76. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/_loop.py +0 -0
  77. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/address.py +0 -0
  78. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/arch.py +0 -0
  79. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/attn_config.py +0 -0
  80. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/attn_scores.py +0 -0
  81. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/autograd_support.py +0 -0
  82. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/chat_compose.py +0 -0
  83. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/chat_conventions.py +0 -0
  84. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/chat_formatters.py +0 -0
  85. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/cuda_preflight.py +0 -0
  86. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/dispatch.py +0 -0
  87. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/hooks.py +0 -0
  88. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/lens.py +0 -0
  89. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/mappers.py +0 -0
  90. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/moe_routing.py +0 -0
  91. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/notebook_stdout.py +0 -0
  92. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/residual_basis.py +0 -0
  93. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/select.py +0 -0
  94. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/steer.py +0 -0
  95. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/steer_specs.py +0 -0
  96. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/sync.py +0 -0
  97. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/tokenize.py +0 -0
  98. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/_demux.py +0 -0
  99. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/_payload.py +0 -0
  100. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/graphs.py +0 -0
  101. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/lens/__init__.py +0 -0
  102. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/lens/intervene.py +0 -0
  103. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/lens/readout.py +0 -0
  104. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/lens/unembed.py +0 -0
  105. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/mhc.py +0 -0
  106. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/native.py +0 -0
  107. {interp_engine-1.7.3 → interp_engine-1.9.0}/interp_engine/vllm_capture/steering.py +0 -0
  108. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/conftest.py +0 -0
  109. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/harness.py +0 -0
  110. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/model_expectations.yaml +0 -0
  111. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/synthetic_families.py +0 -0
  112. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_address.py +0 -0
  113. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_attn_config_tripwire.py +0 -0
  114. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_attn_probs_indexing.py +0 -0
  115. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_attn_scores.py +0 -0
  116. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_attn_z_gqa.py +0 -0
  117. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_autograd_support.py +0 -0
  118. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_bench_workloads.py +0 -0
  119. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_capability_refusals.py +0 -0
  120. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_capture_addressing.py +0 -0
  121. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_chat_compose.py +0 -0
  122. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_chat_formatters.py +0 -0
  123. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_chat_templates.py +0 -0
  124. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_core.py +0 -0
  125. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_cuda_preflight.py +0 -0
  126. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_doc_code_fences.py +0 -0
  127. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_eager_autograd.py +0 -0
  128. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_family_points.py +0 -0
  129. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_gated_attn_out.py +0 -0
  130. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_head_contributions.py +0 -0
  131. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_hook_call_conventions.py +0 -0
  132. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_injected_system_spans.py +0 -0
  133. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_layer_kinds.py +0 -0
  134. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_logit_transform.py +0 -0
  135. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_mappers.py +0 -0
  136. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_mlp_internals.py +0 -0
  137. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_model_expectations.py +0 -0
  138. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_moe.py +0 -0
  139. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_new_models_gpu.py +0 -0
  140. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_no_chat_template.py +0 -0
  141. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_normalized_hook.py +0 -0
  142. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_notebook_stdout.py +0 -0
  143. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_packaging.py +0 -0
  144. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_parity_gpt2.py +0 -0
  145. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_per_layer_attn_dims.py +0 -0
  146. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_points_registry.py +0 -0
  147. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_protocol.py +0 -0
  148. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_published_benchmarks.py +0 -0
  149. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_qk_norm.py +0 -0
  150. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_qkv_layout.py +0 -0
  151. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_reasoning_spans.py +0 -0
  152. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_release.py +0 -0
  153. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_resid_mid.py +0 -0
  154. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_residual_basis.py +0 -0
  155. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_sandwich_norms.py +0 -0
  156. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_select.py +0 -0
  157. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_sliding_window_attn.py +0 -0
  158. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_small_models_gpu.py +0 -0
  159. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_static_dsv4_gpu.py +0 -0
  160. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_static_parity_gpu.py +0 -0
  161. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_static_warmup.py +0 -0
  162. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_steer_context.py +0 -0
  163. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_steer_math_parity.py +0 -0
  164. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_sync_loop.py +0 -0
  165. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_sync_parity.py +0 -0
  166. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_unified_free_functions.py +0 -0
  167. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_unresolved_families.py +0 -0
  168. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_capture_gpu.py +0 -0
  169. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_capture_scales.py +0 -0
  170. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_engine_loop.py +0 -0
  171. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_graph_path.py +0 -0
  172. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_graphs_on_gpu.py +0 -0
  173. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_hook_availability.py +0 -0
  174. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_hyper_connections.py +0 -0
  175. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_only_families.py +0 -0
  176. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_plugin.py +0 -0
  177. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vllm_wire_grammar.py +0 -0
  178. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_vocabulary_boundary.py +0 -0
  179. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_worker_lens_capture_readout.py +0 -0
  180. {interp_engine-1.7.3 → interp_engine-1.9.0}/tests/test_worker_lens_readout.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: interp-engine
3
- Version: 1.7.3
3
+ Version: 1.9.0
4
4
  Summary: A fast, standardized, and easy to use interpretability engine.
5
5
  Project-URL: Homepage, https://github.com/decoderesearch/interp-engine
6
6
  Project-URL: Repository, https://github.com/decoderesearch/interp-engine
@@ -21,7 +21,11 @@ class EnvStamp:
21
21
  """What the numbers were produced on. Recorded per run, so a stale result is identifiable."""
22
22
 
23
23
  gpu_name: str = ""
24
+ """The card, prefixed with the count when the model was sharded across more than one
25
+ (``2x NVIDIA A40``), so every view keyed on this field says how many cards the number is from."""
24
26
  gpu_total_gib: float = 0.0
27
+ """Per card, not summed: the figure a reader compares against a card's spec sheet."""
28
+ gpu_count: int = 1
25
29
  driver_version: str = ""
26
30
  cuda_version: str = ""
27
31
  torch_version: str = ""
@@ -76,11 +80,13 @@ def _driver_version() -> str:
76
80
  return out.stdout.strip().splitlines()[0].strip() if out.stdout.strip() else "unknown"
77
81
 
78
82
 
79
- def env_stamp(index: int = 0) -> EnvStamp:
83
+ def env_stamp(index: int = 0, num_gpus: int = 1) -> EnvStamp:
80
84
  props = torch.cuda.get_device_properties(index)
85
+ name = torch.cuda.get_device_name(index)
81
86
  return EnvStamp(
82
- gpu_name=torch.cuda.get_device_name(index),
87
+ gpu_name=f"{num_gpus}x {name}" if num_gpus > 1 else name,
83
88
  gpu_total_gib=props.total_memory / GIB,
89
+ gpu_count=num_gpus,
84
90
  driver_version=_driver_version(),
85
91
  cuda_version=torch.version.cuda or "unknown",
86
92
  torch_version=torch.__version__,
@@ -11,6 +11,9 @@
11
11
  # bash benchmarks/run_all.sh --models gemma-2-2b,qwen3-4b --variants eager,vllm
12
12
  # bash benchmarks/run_all.sh --workloads generate,capture_mid --no-report
13
13
  # bash benchmarks/run_all.sh --gpu-memory-utilization 0.7 # smaller card, or more worker scratch
14
+ # bash benchmarks/run_all.sh --num-gpus 1 --models qwen3.8-27b # pin a multi-card box to one card
15
+ #
16
+ # Every cell is sharded across every visible CUDA card unless --num-gpus says otherwise.
14
17
  # BENCH_PYTHON=/path/to/venv/bin/python bash benchmarks/run_all.sh
15
18
  #
16
19
  # With no --models, the sweep runs every model in the spec that this card can hold and names the ones
@@ -43,6 +46,7 @@ MODELS=""
43
46
  VARIANTS=""
44
47
  WORKLOADS=""
45
48
  GPU_MEM_UTIL=""
49
+ NUM_GPUS=""
46
50
  RUN_REPORT=1
47
51
  SKIP_EXISTING=0
48
52
 
@@ -52,6 +56,7 @@ while [[ $# -gt 0 ]]; do
52
56
  --variants) VARIANTS="$2"; shift 2 ;;
53
57
  --workloads) WORKLOADS="$2"; shift 2 ;;
54
58
  --gpu-memory-utilization) GPU_MEM_UTIL="$2"; shift 2 ;;
59
+ --num-gpus) NUM_GPUS="$2"; shift 2 ;;
55
60
  --python) PYTHON="$2"; shift 2 ;;
56
61
  --no-report) RUN_REPORT=0; shift ;;
57
62
  # Resume a sweep that was interrupted. Off by default: a normal rerun should replace stale
@@ -70,13 +75,30 @@ if ! "$PYTHON" -c 'import interp_engine' 2>/dev/null; then
70
75
  exit 2
71
76
  fi
72
77
 
73
- # Total VRAM on the card this sweep will use, in GiB, or empty if there is no nvidia-smi to ask.
74
- # GiB rather than the vendor's GB, to match `min_gpu_gib` in the spec and the `gpu_total_gib` every
75
- # cell records -- a "180 GB" B200 reads as 179 GiB here, and comparing the two units is how a row
76
- # gets dropped on the one card that fits it.
78
+ # No --num-gpus means every card the process can see: CUDA_VISIBLE_DEVICES when set, else what
79
+ # nvidia-smi lists, else 1.
80
+ if [[ -z "$NUM_GPUS" ]]; then
81
+ if [[ -n "${CUDA_VISIBLE_DEVICES+x}" ]]; then
82
+ NUM_GPUS=$(echo "$CUDA_VISIBLE_DEVICES" | tr ',' '\n' | grep -c .)
83
+ else
84
+ NUM_GPUS=$(nvidia-smi -L 2>/dev/null | grep -c '^GPU')
85
+ fi
86
+ (( NUM_GPUS > 0 )) || NUM_GPUS=1
87
+ fi
88
+ if ! [[ "$NUM_GPUS" =~ ^[1-9][0-9]*$ ]]; then
89
+ echo "error: --num-gpus must be a positive integer (got '$NUM_GPUS')" >&2
90
+ exit 2
91
+ fi
92
+
93
+ # Total VRAM this sweep will use, in GiB, or empty if there is no nvidia-smi to ask: the first
94
+ # --num-gpus cards summed, since that is what a sharded load has to fit into. GiB rather than the
95
+ # vendor's GB, to match `min_gpu_gib` in the spec and the `gpu_total_gib` every cell records -- a
96
+ # "180 GB" B200 reads as 179 GiB here, and comparing the two units is how a row gets dropped on the
97
+ # one card that fits it.
77
98
  gpu_total_gib() {
78
99
  local mib
79
- mib="$(nvidia-smi --query-gpu=memory.total --format=csv,noheader,nounits 2>/dev/null | head -1 | tr -d ' ')"
100
+ mib="$(nvidia-smi --query-gpu=memory.total --format=csv,noheader,nounits 2>/dev/null \
101
+ | head -n "$NUM_GPUS" | tr -d ' ' | awk '{ s += $1 } END { if (NR) print s }')"
80
102
  [[ -z "$mib" ]] && return 0
81
103
  awk -v mib="$mib" 'BEGIN { printf "%.1f", mib / 1024 }'
82
104
  }
@@ -171,6 +193,7 @@ for pair in "${PAIRS[@]}"; do
171
193
  args=(--model "$model" --variant "$variant")
172
194
  [[ -n "$WORKLOADS" ]] && args+=(--workloads "$WORKLOADS")
173
195
  [[ -n "$GPU_MEM_UTIL" ]] && args+=(--gpu-memory-utilization "$GPU_MEM_UTIL")
196
+ (( NUM_GPUS > 1 )) && args+=(--num-gpus "$NUM_GPUS")
174
197
 
175
198
  # TOKENIZERS_PARALLELISM: the tokenizer is forked by vLLM's workers and warns on every cell
176
199
  # otherwise. VLLM_LOGGING_LEVEL: vLLM's per-step INFO logging would bury the workload lines.
@@ -200,6 +223,7 @@ if (( RUN_REPORT )); then
200
223
  [[ -n "$VARIANTS" ]] && sweep_cmd="$sweep_cmd --variants $VARIANTS"
201
224
  [[ -n "$WORKLOADS" ]] && sweep_cmd="$sweep_cmd --workloads $WORKLOADS"
202
225
  [[ -n "$GPU_MEM_UTIL" ]] && sweep_cmd="$sweep_cmd --gpu-memory-utilization $GPU_MEM_UTIL"
226
+ (( NUM_GPUS > 1 )) && sweep_cmd="$sweep_cmd --num-gpus $NUM_GPUS"
203
227
  "$PYTHON" -m benchmarks.report_bench --sweep-command "$sweep_cmd"
204
228
  fi
205
229
 
@@ -209,9 +209,13 @@ def _load_kwargs(
209
209
  m: ModelSpec,
210
210
  *,
211
211
  gpu_memory_utilization: float = GPU_MEMORY_UTILIZATION,
212
+ num_gpus: int = 1,
212
213
  ) -> dict[str, Any]:
213
214
  """Variant kwargs plus the per-backend knobs that only one backend accepts.
214
215
 
216
+ ``num_gpus`` is ``load_model``'s own knob and goes through as is: tensor parallelism on vLLM,
217
+ accelerate's layer placement on eager, where it takes the place of ``device="cuda"``.
218
+
215
219
  ``max_model_len`` and ``gpu_memory_utilization`` are vLLM-only, and passing them to the eager
216
220
  constructor would raise. Kept here rather than duplicated into every vLLM variant so the variant
217
221
  table stays about the thing being varied.
@@ -234,6 +238,8 @@ def _load_kwargs(
234
238
  kwargs = dict(v.kwargs)
235
239
  if kwargs.get("static_writes") == STEER_WRITES:
236
240
  kwargs["static_writes"] = [_steer_site(v, m)]
241
+ if num_gpus > 1:
242
+ kwargs["num_gpus"] = num_gpus
237
243
  # Every vLLM backend, not just the hooked one: these are engine settings, and a graph variant that
238
244
  # silently lost `max_model_len` would be measured against a different context than the row beside
239
245
  # it. `VLLM_BACKENDS` is the engine's own list, so a fourth backend is covered the day it lands.
@@ -270,9 +276,10 @@ def _load_kwargs(
270
276
  if isinstance(declared_model_kwargs, dict):
271
277
  merged.update(declared_model_kwargs)
272
278
  kwargs["model_kwargs"] = merged
273
- # A device_map places the weights itself, and `load_model` drops `device` when one is given.
274
- # Setting it anyway would put a `device="cuda"` in the recorded kwargs that had no effect.
275
- if "device_map" not in kwargs:
279
+ # A device_map places the weights itself, and `load_model` drops `device` when one is given
280
+ # -- as it does at `num_gpus > 1`, which is a device_map. Setting it anyway would put a
281
+ # `device="cuda"` in the recorded kwargs that had no effect.
282
+ if "device_map" not in kwargs and num_gpus == 1:
276
283
  kwargs.setdefault("device", "cuda")
277
284
  return kwargs
278
285
 
@@ -284,10 +291,11 @@ async def run_cell(
284
291
  *,
285
292
  command: str,
286
293
  gpu_memory_utilization: float = GPU_MEMORY_UTILIZATION,
294
+ num_gpus: int = 1,
287
295
  ) -> dict[str, Any]:
288
296
  from interp_engine import load_model
289
297
 
290
- stamp = env_stamp()
298
+ stamp = env_stamp(num_gpus=num_gpus)
291
299
  record: dict[str, Any] = {
292
300
  "schema": SCHEMA,
293
301
  "started_at": datetime.now(UTC).isoformat(timespec="seconds"),
@@ -298,7 +306,7 @@ async def run_cell(
298
306
  "env": dataclasses.asdict(stamp),
299
307
  "workloads": {},
300
308
  }
301
- kwargs = _load_kwargs(variant, model_spec, gpu_memory_utilization=gpu_memory_utilization)
309
+ kwargs = _load_kwargs(variant, model_spec, gpu_memory_utilization=gpu_memory_utilization, num_gpus=num_gpus)
302
310
  record["variant"]["kwargs"] = {k: str(val) for k, val in kwargs.items()}
303
311
  record["model"]["native_dtype"] = _native_dtype(model_spec.hf_id)
304
312
 
@@ -392,8 +400,18 @@ def _parse_args(argv: list[str] | None = None) -> argparse.Namespace:
392
400
  f"model's own declared fraction, or {GPU_MEMORY_UTILIZATION} where it has none"
393
401
  ),
394
402
  )
403
+ p.add_argument(
404
+ "--num-gpus",
405
+ type=int,
406
+ default=1,
407
+ help="shard the model across this many cards (vLLM tensor parallelism; accelerate placement "
408
+ "on eager). Stamped on the cell, since a two-card number is not a one-card number",
409
+ )
395
410
  p.add_argument("--list", action="store_true", help="print the models, variants and workloads and exit")
396
- return p.parse_args(argv)
411
+ args = p.parse_args(argv)
412
+ if args.num_gpus < 1:
413
+ p.error("--num-gpus must be at least 1")
414
+ return args
397
415
 
398
416
 
399
417
  def main(argv: list[str] | None = None) -> int:
@@ -455,6 +473,7 @@ def main(argv: list[str] | None = None) -> int:
455
473
  workload_keys,
456
474
  command=command,
457
475
  gpu_memory_utilization=bench_spec.gpu_memory_utilization_for(model_spec.key, args.gpu_memory_utilization),
476
+ num_gpus=args.num_gpus,
458
477
  )
459
478
  )
460
479
 
@@ -180,7 +180,7 @@ Canonical names, with the layer after a dot: `resid_post.10`. Extra coordinates
180
180
  | `mlp_act`, `mlp_pre`, `mlp_pre_linear` | MLP internals, `d_mlp` wide | **eager only** |
181
181
  | `router_logits` | MoE routing scores, every expert | both |
182
182
  | `expert_weights`, `expert_indices` | the top-k the router selected, and its weights | **eager only** |
183
- | the QK-norm points | inside the attention module | both, single-GPU only (head-sharded) |
183
+ | the QK-norm points | inside the attention module | both (head-sharded; gathered across TP ranks) |
184
184
 
185
185
  Eager-only is not an omission: vLLM's fused MLP and MoE kernels compute those tensors inline, so
186
186
  there is no module boundary to hook. This table is the working subset;
@@ -191,7 +191,8 @@ Attention is the one row that reads "both" with a caveat. No boundary holds a sc
191
191
  backend, so `capture_attention(model, tokens, layers)` is how you ask, and it returns the same
192
192
  `{layer: {"scores", "probs", "value"}}` either way — from `output_attentions` on eager (which needs
193
193
  the model loaded with `attn_implementation="eager"`) and from an off-kernel recompute over captured
194
- post-RoPE q/k on vLLM (single-GPU only). Different code paths, same contract; `value` there is the
194
+ post-RoPE q/k on vLLM (gathered across ranks under tensor parallelism, so it works at any
195
+ `num_gpus`). Different code paths, same contract; `value` there is the
195
196
  per-head, family-scaled tensor satisfying `probs @ value == z`, not the raw projection output.
196
197
 
197
198
  [ENGINE_HOOK_MAPPINGS.md](ENGINE_HOOK_MAPPINGS.md) is the full dictionary across all three
@@ -228,9 +229,10 @@ an obvious one.
228
229
  always, on every configuration — the unembed happens in another process. Eager can do it, but only
229
230
  if the model was loaded with `requires_grad=True`. Gate on `model.grad_support`, not on backend
230
231
  name. See [GRADIENTS.md](GRADIENTS.md).
231
- 5. **vLLM with `num_gpus > 1` serves no `z`, no DFA and no attention recompute.** Heads are sharded
232
- across ranks, so rank 0 holds a slice; the engine refuses rather than returning it. Use one GPU, or
233
- the eager backend, for per-head work.
232
+ 5. **vLLM with `num_gpus > 1` serves the same points as one GPU.** Heads and MLP neurons are
233
+ sharded across ranks, and the worker gathers `z`, `value`, `mlp_act`, the QK-norm points and the
234
+ q/k behind `capture_attention` back to full width at collect, so rank 0's payload is the whole
235
+ tensor. Every `collect_*` is a collective on every rank; do not call one from a single rank.
234
236
  6. **`await model.warmup()` before timing anything.** Construction is deliberately cheap and lazy on
235
237
  both backends — on vLLM nearly the whole load happens in `warmup()`, so without it your first
236
238
  request's latency is the load time.
@@ -108,7 +108,15 @@ index names the layer HF's does, which nothing checked before and which would fa
108
108
  than raise. It needs `interp-engine[vllm]`, so it self-skips elsewhere; note that running it via
109
109
  `.venv-vllm/bin/python` needs that directory on `PATH` too, because vLLM shells out to `ninja` to
110
110
  build a sampler kernel at startup. `tests/test_vllm_wire_grammar.py` covers the same process
111
- boundary on CPU, over a synthetic demux.
111
+ boundary on CPU, over a synthetic demux. `tests/test_multigpu.py` (`-m multigpu`, two CUDA cards)
112
+ repeats the comparison at `num_gpus=2`: eager under accelerate's layer placement and vLLM under
113
+ tensor parallelism, where `vllm_capture/_tp.py` gathers the head- and neuron-sharded points across
114
+ ranks, against the one-card eager reference — captures, attention, steering, decode and the lens.
115
+ One rule follows from tensor parallelism: a worker RPC refuses by *returning* a reason, never by
116
+ raising. Under vLLM's multiprocess executor a raise inside `collective_rpc` is read from one rank
117
+ while the others' replies stay queued, and the next RPC consumes those — so a single refused layer
118
+ took every point of a TP=8 cell with it. `resolvable_points` and `resolvable_attn` are the ask-first
119
+ calls; `tests/test_vllm_new_points.py` pins that they install nothing.
112
120
 
113
121
  **The docs are parsed, not maintained.** The point table in
114
122
  [SUPPORTED_POINTS.md](SUPPORTED_POINTS.md) and the footnote markers in
@@ -14,12 +14,12 @@ nnsight and nnterp.
14
14
  | [`embeddings`][embeddings] | `d_model` | ✅ | ✅ | trunk-level, so addressed with no layer index; distinct from `resid_pre` at layer 0 only where the trunk adds positional embeddings or scales the embedding |
15
15
  | [`resid_pre`][resid_pre] | `d_model` | ✅ | ✅ | |
16
16
  | [`attn_in`][attn_in] | `d_model` | ✅ | ✅ | |
17
- | [`q_norm_in`][q_norm_in] / [`q_norm_out`][q_norm_out] | `n_heads * head_dim` | ✅ | ✅ | head-sharded, so single-GPU only |
18
- | [`k_norm_in`][k_norm_in] / [`k_norm_out`][k_norm_out] | `n_kv_heads * head_dim` | ✅ | ✅ | head-sharded, so single-GPU only |
19
- | [`value`][value] | `n_heads * head_dim` | ✅ | ✅ | head-sharded, so single-GPU only |
17
+ | [`q_norm_in`][q_norm_in] / [`q_norm_out`][q_norm_out] | `n_heads * head_dim` | ✅ | ✅ | head-sharded; gathered across TP ranks at collect |
18
+ | [`k_norm_in`][k_norm_in] / [`k_norm_out`][k_norm_out] | `n_kv_heads * head_dim` | ✅ | ✅ | head-sharded; gathered across TP ranks at collect |
19
+ | [`value`][value] | `n_heads * head_dim` | ✅ | ✅ | head-sharded; gathered across TP ranks at collect |
20
20
  | [`attn_scores`][attn_scores] | `n_heads * query * key` | ✅ | ♻️ | no module boundary holds the pre-softmax matrix on **either** backend; vLLM rebuilds it from captured post-RoPE q/k |
21
21
  | [`attn_probs`][attn_probs] | `n_heads * query * key` | ✅ | ♻️ | fused paged attention never materializes the probabilities; same recompute |
22
- | [`z`][z] | `n_heads * head_dim` | ✅ | ✅ | head-sharded, so single-GPU only |
22
+ | [`z`][z] | `n_heads * head_dim` | ✅ | ✅ | head-sharded; gathered across TP ranks at collect |
23
23
  | [`attn_gate`][attn_gate] | `n_heads * head_dim` | ✅ | ❌ | unimplemented — a real module on both trees |
24
24
  | [`attn_out`][attn_out] | `d_model` | ✅ | ✅ | |
25
25
  | [`attn_out_post`][attn_out_post] | `d_model` | ✅ | ✅ | |
@@ -27,7 +27,7 @@ nnsight and nnterp.
27
27
  | [`mlp_in`][mlp_in] | `d_model` | ✅ | ✅ | |
28
28
  | [`mlp_pre`][mlp_pre] | `d_mlp` | ✅ | ❌ | unreachable — vLLM fuses `gate_proj` and `up_proj` into one `gate_up_proj`, so neither branch is a module output |
29
29
  | [`mlp_pre_linear`][mlp_pre_linear] | `d_mlp` | ✅ | ❌ | as `mlp_pre`; gated MLPs only |
30
- | [`mlp_act`][mlp_act] | `d_mlp` | ✅ | ✅ | neuron-sharded, so single-GPU only |
30
+ | [`mlp_act`][mlp_act] | `d_mlp` | ✅ | ✅ | neuron-sharded; gathered across TP ranks at collect |
31
31
  | [`router_logits`][router_logits] | `n_experts` | ✅ | ✅ | replicated gate, so it survives tensor parallelism |
32
32
  | [`expert_weights`][expert_weights] | `n_experts` | ✅ | ❌ | unreachable — the top-k happens inside the FusedMoE kernel, which returns the combined output with the selection never materialized |
33
33
  | [`expert_indices`][expert_indices] | `n_experts` | ✅ | ❌ | as `expert_weights` |
@@ -54,13 +54,16 @@ a fused kernel ate the tensor and no module boundary holds it. Ask the code rath
54
54
  you are branching on it — `points.vllm_hookable()` is the served set, `points.reason(name)` is the
55
55
  sentence for one refusal, and `model.points()` is what a loaded model has.
56
56
 
57
- ## Tensor parallelism narrows the vLLM column further
57
+ ## Tensor parallelism does not narrow the vLLM column
58
58
 
59
- The capture path reads rank 0's payload alone, so a point whose last axis vLLM shards comes back as a
60
- slice: `z`, `value`, `mlp_act` and the four QK-norm points are refused on a multi-GPU pod rather than
61
- returned short, and so is the attention recompute (q/k/v are head-sharded). Everything `d_model` wide
62
- is all-reduced before the hook sees it, and `router_logits` comes off a replicated gate, so those are
63
- unaffected.
59
+ The capture path reads rank 0's payload alone, and at `num_gpus > 1` vLLM shards `z`, `value`,
60
+ `mlp_act`, the four QK-norm points and the q/k/v behind the attention recompute by head or by neuron.
61
+ The worker gathers those across ranks at collect time (`interp_engine.vllm_capture._tp`), so rank 0
62
+ hands back the same full-width tensor a single card would have, in the same layout; a KV head that
63
+ vLLM replicates rather than shards is kept once. Everything `d_model` wide is all-reduced before the
64
+ hook sees it, and `router_logits` comes off a replicated gate, so those never needed gathering.
65
+ Verified against the eager reference at TP=2 by the validator (`validator/`, `NUM_GPUS=2`) and by
66
+ `tests/test_multigpu.py`.
64
67
 
65
68
  ## The last seven rows need a hyper-connection trunk
66
69
 
@@ -94,6 +94,16 @@ eager = load_model("google/gemma-2-2b-it", backend="eager", device="cuda", dtype
94
94
  served = load_model("meta-llama/Llama-3.1-8B", backend="vllm", gpu_memory_utilization=0.85)
95
95
  ```
96
96
 
97
+ Two knobs narrow a bf16 checkpoint on load, with no calibration and no other repo:
98
+ `quantization=` applies `"fp8"` (vLLM backends), `"bnb-4bit"` (either backend) or `"bnb-8bit"`
99
+ (eager) to every linear layer, and `kv_cache_dtype="fp8"` halves vLLM's paged cache. Both take the
100
+ same name on both backends and are refused, naming the alternative, where a backend cannot apply
101
+ them. The [GPU sizer](../gpu-sizer/INPUTS.md) prices both and prints the argument it priced.
102
+
103
+ ```python
104
+ narrow = load_model("meta-llama/Llama-3.3-70B-Instruct", backend="vllm", quantization="fp8", kv_cache_dtype="fp8")
105
+ ```
106
+
97
107
  **Construction is cheap and lazy on both backends; `warmup()` is where the cost lands.** On vLLM
98
108
  almost the entire load happens there, so call it before you time anything or a first request pays
99
109
  for the engine.
@@ -480,8 +480,8 @@ def capture_attention(
480
480
  pass rather than one rebuilt from the other.
481
481
 
482
482
  The eager arm needs the model loaded with eager attention for ``attn_probs``; the vLLM arm
483
- recomputes off-kernel from captured post-RoPE q/k and is single-GPU only. Both refusals name
484
- themselves.
483
+ recomputes off-kernel from captured post-RoPE q/k, gathered across ranks under tensor
484
+ parallelism. Both refusals name themselves.
485
485
  """
486
486
  if not isinstance(model, EagerModel):
487
487
  ids = as_token_ids(tokens, model=model, what="capture_attention")
@@ -1630,6 +1630,21 @@ def value_head_dim(cfg: Any, head_dim: int) -> int:
1630
1630
  return _first_int(cfg, ("v_head_dim",)) or head_dim
1631
1631
 
1632
1632
 
1633
+ def kv_latent_width(cfg: Any) -> int:
1634
+ """Elements one token of KV cache holds per layer under MLA, or 0 where keys and values are cached.
1635
+
1636
+ DeepSeek's multi-head latent attention never caches K and V: it caches the ``kv_lora_rank``-wide
1637
+ latent both are expanded from, plus the ``qk_rope_head_dim`` positional key part that rides beside
1638
+ it -- one 576-wide row per token on DeepSeek-V3 and Kimi-K2, in place of 64 heads of K and V. A
1639
+ cache sized from ``n_kv_heads * (head_dim + v_head_dim)`` is 21x too large there, and it is the
1640
+ figure vLLM builds its pool from, so this is what a sizer has to charge. 0 on every non-MLA config.
1641
+ """
1642
+ rank = _first_int(cfg, ("kv_lora_rank",))
1643
+ if not rank:
1644
+ return 0
1645
+ return rank + (_first_int(cfg, ("qk_rope_head_dim",)) or 0)
1646
+
1647
+
1633
1648
  def first_kv_shared_layer(cfg: Any, n_layers: int) -> int | None:
1634
1649
  """The first layer that reuses an earlier layer's keys/values, or None if none do."""
1635
1650
  shared = _first_int(cfg, ("num_kv_shared_layers",))
@@ -2020,6 +2035,9 @@ class ModelFacts:
2020
2035
  # The width of one value head, which differs from ``head_dim`` on MiMo-V2 and the DeepSeek MLA
2021
2036
  # families. See :func:`value_head_dim`; ``value`` and ``z`` are this wide per head, not ``head_dim``.
2022
2037
  v_head_dim: int = 0
2038
+ # Elements per token per layer that an MLA trunk caches in place of K and V (the latent plus the
2039
+ # RoPE'd key part); 0 where the cache holds K and V. See :func:`kv_latent_width`.
2040
+ kv_latent_width: int = 0
2023
2041
  # First layer that reuses an earlier layer's keys/values and so has no k/v projection of its own
2024
2042
  # (Gemma-4); None when every layer computes its own.
2025
2043
  first_kv_shared_layer: int | None = None
@@ -2476,6 +2494,7 @@ def resolve_facts(config: Any, *, n_layers_fallback: int | None = None) -> Model
2476
2494
  global_kv_heads=_first_int(cfg, ("num_global_key_value_heads",)) or None,
2477
2495
  k_eq_v=bool(config_attr(cfg, "attention_k_eq_v", False)),
2478
2496
  v_head_dim=value_head_dim(cfg, head_dim),
2497
+ kv_latent_width=kv_latent_width(cfg),
2479
2498
  first_kv_shared_layer=first_kv_shared_layer(cfg, n_layers),
2480
2499
  n_experts=n_experts(cfg),
2481
2500
  experts_per_token=_first_int(cfg, _EXPERTS_PER_TOKEN_FIELDS) or 0,
@@ -50,12 +50,62 @@ def _declares_nothing(value: Any) -> bool:
50
50
  return not list(value)
51
51
 
52
52
 
53
+ def _apply_load_precision(backend: str, quantization: str, kv_cache_dtype: str, backend_kwargs: dict[str, Any]) -> None:
54
+ """Turn ``quantization`` and ``kv_cache_dtype`` into what ``backend``'s constructor takes.
55
+
56
+ Both are one name on ``load_model`` and two different things underneath: vLLM quantizes through
57
+ an engine argument, transformers through a ``BitsAndBytesConfig``. The table in
58
+ :data:`interp_engine.memory.QUANTIZATIONS` says which backend applies which scheme, and a scheme
59
+ the backend cannot apply is refused here with that table's reason -- never passed on to become an
60
+ opaque ``TypeError`` from a constructor, and never dropped to load the checkpoint as stored.
61
+ """
62
+ from interp_engine.memory import QUANTIZATIONS, quantization_refusal
63
+
64
+ use_vllm = backend in VLLM_BACKENDS
65
+ if kv_cache_dtype not in ("auto", "", None):
66
+ if not use_vllm:
67
+ raise ValueError(
68
+ f"kv_cache_dtype={kv_cache_dtype!r} names the dtype of vLLM's paged KV cache, and "
69
+ f"backend={backend!r} has no such cache. Drop it, or use a vLLM backend."
70
+ )
71
+ extra = dict(backend_kwargs.get("extra_vllm_kwargs") or {})
72
+ extra.setdefault("kv_cache_dtype", kv_cache_dtype)
73
+ backend_kwargs["extra_vllm_kwargs"] = extra
74
+
75
+ if not quantization:
76
+ return
77
+ refused = quantization_refusal(quantization, backend)
78
+ if refused:
79
+ raise ValueError(f"quantization={quantization!r} on backend={backend!r}: {refused}")
80
+ scheme = QUANTIZATIONS[quantization]
81
+ if use_vllm:
82
+ extra = dict(backend_kwargs.get("extra_vllm_kwargs") or {})
83
+ if extra.get("quantization") not in (None, scheme.vllm_name):
84
+ raise ValueError(
85
+ f"quantization={quantization!r} asks vLLM for {scheme.vllm_name!r}, but extra_vllm_kwargs "
86
+ f"already names {extra['quantization']!r}. Pass one or the other."
87
+ )
88
+ extra["quantization"] = scheme.vllm_name
89
+ backend_kwargs["extra_vllm_kwargs"] = extra
90
+ return
91
+ if backend_kwargs.get("quantization_config") is not None:
92
+ raise ValueError(
93
+ f"quantization={quantization!r} builds a BitsAndBytesConfig, and quantization_config= was "
94
+ f"passed as well. Pass one or the other."
95
+ )
96
+ from transformers import BitsAndBytesConfig
97
+
98
+ backend_kwargs["quantization_config"] = BitsAndBytesConfig(**scheme.eager_config)
99
+
100
+
53
101
  def load_model(
54
102
  hf_model_id: str,
55
103
  *,
56
104
  backend: str = "auto",
57
105
  device: str | None = None,
58
106
  dtype: str = "auto",
107
+ quantization: str = "",
108
+ kv_cache_dtype: str = "auto",
59
109
  num_gpus: int = 1,
60
110
  trust_remote_code: bool | None = None,
61
111
  static_points: Any = None,
@@ -86,11 +136,26 @@ def load_model(
86
136
  device: Explicit device for the eager backend. None means let the ladder choose.
87
137
  Ignored by vLLM, which always initializes on CUDA.
88
138
  dtype: ``"auto"`` (the checkpoint's native precision) or an explicit
89
- ``"float32"``/``"float16"``/``"bfloat16"``.
139
+ ``"float32"``/``"float16"``/``"bfloat16"``. This is the width the **activations** run
140
+ at, and the width an unquantized checkpoint's weights are held at. It is not how to ask
141
+ for a narrower checkpoint: vLLM rejects ``dtype="fp8"``, and transformers would store
142
+ fp8 weights with no kernel behind them. That is ``quantization``.
143
+ quantization: An on-load scheme from :data:`interp_engine.memory.QUANTIZATIONS`, applied to
144
+ a wider checkpoint as it loads, with no calibration step: ``"fp8"`` (vLLM only),
145
+ ``"bnb-4bit"`` (either backend) or ``"bnb-8bit"`` (eager only). Empty, the default,
146
+ loads the checkpoint as stored -- which for a repo that already ships quantized is the
147
+ right answer, since a quantizer cannot narrow what is already narrower. A scheme the
148
+ chosen backend cannot apply is refused, naming the one to use instead. Other vLLM
149
+ schemes still reach the engine through ``extra_vllm_kwargs={"quantization": ...}``.
150
+ kv_cache_dtype: vLLM's KV cache dtype -- ``"auto"`` (the model dtype, or the scheme the
151
+ checkpoint declares for its cache) or ``"fp8"``, which halves the cache and so roughly
152
+ doubles the context or concurrency a card holds. Refused on the eager backend, which
153
+ has no paged cache to set the dtype of.
90
154
  num_gpus: Shard across this many GPUs on one node -- vLLM ``tensor_parallel_size``,
91
- eager accelerate ``device_map="auto"``. Note that vLLM with ``num_gpus > 1``
92
- cannot serve per-head ``z`` or DFA, because attention heads are sharded across
93
- ranks and the off-kernel recompute would only see one shard.
155
+ eager accelerate ``device_map="auto"``. The vLLM worker gathers the head- and
156
+ neuron-sharded points (``z``, ``value``, ``mlp_act``, the QK-norm points, the q/k
157
+ behind the attention recompute) across ranks at collect, so the served point set
158
+ and every tensor's width are the same as on one GPU.
94
159
  trust_remote_code: Passed to both the config probe and the backend. The default ``None``
95
160
  means "only where the checkpoint has no alternative": eager prefers a native
96
161
  transformers class over a checkpoint's bundled copy of one when both exist, since the
@@ -118,7 +183,8 @@ def load_model(
118
183
  ValueError: ``backend`` is not one of :data:`BACKENDS`; or ``static_points`` /
119
184
  ``static_writes`` was passed on a backend other than ``"vllm-static"``; or
120
185
  ``backend="vllm-static"`` declared no taps at all; or ``enforce_eager=True`` was
121
- passed alongside a graph-replaying backend.
186
+ passed alongside a graph-replaying backend; or ``quantization`` / ``kv_cache_dtype``
187
+ asks the chosen backend for something it cannot apply.
122
188
  RuntimeError: a vLLM backend was requested but vLLM is not installed.
123
189
  GradientsUnsupported: ``requires_grad=True`` on a vLLM backend, which cannot
124
190
  provide gradients through its forward on any configuration.
@@ -182,6 +248,8 @@ def load_model(
182
248
  # installs nothing in Worker.load_model, and leaves hooks_available False.
183
249
  static_points, static_writes = [], None
184
250
 
251
+ _apply_load_precision(resolved, quantization, kv_cache_dtype, backend_kwargs)
252
+
185
253
  if use_vllm:
186
254
  require_vllm(f"backend={resolved!r} requested for {hf_model_id}")
187
255
  # `requires_grad` is an eager-only constructor kwarg, so on vLLM it would otherwise land as