interp-engine 1.7.2__tar.gz → 1.8.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 (178) hide show
  1. {interp_engine-1.7.2 → interp_engine-1.8.0}/PKG-INFO +1 -1
  2. {interp_engine-1.7.2 → interp_engine-1.8.0}/docs/USAGE.md +13 -1
  3. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/load.py +69 -2
  4. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/memory.py +246 -16
  5. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/tokenize.py +78 -2
  6. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_backend.py +12 -5
  7. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/static.py +28 -0
  8. {interp_engine-1.7.2 → interp_engine-1.8.0}/pyproject.toml +1 -1
  9. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_gpu_sizer.py +41 -6
  10. interp_engine-1.8.0/tests/test_injected_system_spans.py +115 -0
  11. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_load.py +81 -0
  12. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_memory.py +178 -1
  13. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_static_set.py +49 -0
  14. {interp_engine-1.7.2 → interp_engine-1.8.0}/.gitignore +0 -0
  15. {interp_engine-1.7.2 → interp_engine-1.8.0}/LICENSE +0 -0
  16. {interp_engine-1.7.2 → interp_engine-1.8.0}/README.md +0 -0
  17. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/README.md +0 -0
  18. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/__init__.py +0 -0
  19. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/bench_spec.py +0 -0
  20. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/cells.py +0 -0
  21. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/probe.py +0 -0
  22. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/probe_lens_stream.py +0 -0
  23. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/publish.py +0 -0
  24. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/report_bench.py +0 -0
  25. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/deepseek-v4-flash-0731__eager.json +0 -0
  26. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/deepseek-v4-flash-0731__vllm-cudagraph.json +0 -0
  27. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/deepseek-v4-flash-0731__vllm-static.json +0 -0
  28. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/deepseek-v4-flash-0731__vllm.json +0 -0
  29. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/gemma-2-2b__eager.json +0 -0
  30. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/gemma-2-2b__vllm-cudagraph.json +0 -0
  31. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/gemma-2-2b__vllm-static.json +0 -0
  32. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/gemma-2-2b__vllm.json +0 -0
  33. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/llama-3.1-8b__eager.json +0 -0
  34. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/llama-3.1-8b__vllm-cudagraph.json +0 -0
  35. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/llama-3.1-8b__vllm-static.json +0 -0
  36. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/llama-3.1-8b__vllm.json +0 -0
  37. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/qwen3-4b__eager.json +0 -0
  38. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/qwen3-4b__vllm-cudagraph.json +0 -0
  39. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/qwen3-4b__vllm-static.json +0 -0
  40. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/qwen3-4b__vllm.json +0 -0
  41. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/qwen3.8-27b__eager.json +0 -0
  42. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/qwen3.8-27b__vllm-cudagraph.json +0 -0
  43. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/qwen3.8-27b__vllm-static.json +0 -0
  44. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results/qwen3.8-27b__vllm.json +0 -0
  45. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/results-latest.md +0 -0
  46. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/run_all.sh +0 -0
  47. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/run_bench.py +0 -0
  48. {interp_engine-1.7.2 → interp_engine-1.8.0}/benchmarks/workloads.py +0 -0
  49. {interp_engine-1.7.2 → interp_engine-1.8.0}/docs/AGENT_INTEGRATION.md +0 -0
  50. {interp_engine-1.7.2 → interp_engine-1.8.0}/docs/ARCHITECTURE_QUIRKS.md +0 -0
  51. {interp_engine-1.7.2 → interp_engine-1.8.0}/docs/COMPATIBILITY.md +0 -0
  52. {interp_engine-1.7.2 → interp_engine-1.8.0}/docs/ENGINE_HOOK_MAPPINGS.md +0 -0
  53. {interp_engine-1.7.2 → interp_engine-1.8.0}/docs/GRADIENTS.md +0 -0
  54. {interp_engine-1.7.2 → interp_engine-1.8.0}/docs/INTERNALS.md +0 -0
  55. {interp_engine-1.7.2 → interp_engine-1.8.0}/docs/PERFORMANCE.md +0 -0
  56. {interp_engine-1.7.2 → interp_engine-1.8.0}/docs/PORTING.md +0 -0
  57. {interp_engine-1.7.2 → interp_engine-1.8.0}/docs/README.md +0 -0
  58. {interp_engine-1.7.2 → interp_engine-1.8.0}/docs/SUPPORTED_POINTS.md +0 -0
  59. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/__init__.py +0 -0
  60. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/_loop.py +0 -0
  61. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/address.py +0 -0
  62. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/arch.py +0 -0
  63. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/attn_config.py +0 -0
  64. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/attn_scores.py +0 -0
  65. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/autograd_support.py +0 -0
  66. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/capture.py +0 -0
  67. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/chat_compose.py +0 -0
  68. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/chat_conventions.py +0 -0
  69. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/chat_formatters.py +0 -0
  70. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/cuda_preflight.py +0 -0
  71. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/dispatch.py +0 -0
  72. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/facts.py +0 -0
  73. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/hooks.py +0 -0
  74. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/lens.py +0 -0
  75. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/mappers.py +0 -0
  76. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/model.py +0 -0
  77. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/moe_routing.py +0 -0
  78. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/notebook_stdout.py +0 -0
  79. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/points.py +0 -0
  80. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/protocol.py +0 -0
  81. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/residual_basis.py +0 -0
  82. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/select.py +0 -0
  83. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/steer.py +0 -0
  84. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/steer_specs.py +0 -0
  85. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/sync.py +0 -0
  86. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/__init__.py +0 -0
  87. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/_demux.py +0 -0
  88. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/_hooks.py +0 -0
  89. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/_payload.py +0 -0
  90. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/_tree.py +0 -0
  91. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/attn.py +0 -0
  92. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/capture.py +0 -0
  93. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/graphs.py +0 -0
  94. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/lens/__init__.py +0 -0
  95. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/lens/intervene.py +0 -0
  96. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/lens/readout.py +0 -0
  97. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/lens/unembed.py +0 -0
  98. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/mhc.py +0 -0
  99. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/native.py +0 -0
  100. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/requests.py +0 -0
  101. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_capture/steering.py +0 -0
  102. {interp_engine-1.7.2 → interp_engine-1.8.0}/interp_engine/vllm_plugin.py +0 -0
  103. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/conftest.py +0 -0
  104. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/harness.py +0 -0
  105. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/model_expectations.yaml +0 -0
  106. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/synthetic_families.py +0 -0
  107. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_address.py +0 -0
  108. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_attn_config_tripwire.py +0 -0
  109. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_attn_probs_indexing.py +0 -0
  110. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_attn_scores.py +0 -0
  111. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_attn_z_gqa.py +0 -0
  112. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_autograd_support.py +0 -0
  113. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_bench_workloads.py +0 -0
  114. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_capability_refusals.py +0 -0
  115. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_capture_addressing.py +0 -0
  116. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_chat_compose.py +0 -0
  117. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_chat_formatters.py +0 -0
  118. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_chat_templates.py +0 -0
  119. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_core.py +0 -0
  120. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_cuda_preflight.py +0 -0
  121. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_doc_code_fences.py +0 -0
  122. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_eager_autograd.py +0 -0
  123. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_facts.py +0 -0
  124. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_family_points.py +0 -0
  125. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_gated_attn_out.py +0 -0
  126. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_head_contributions.py +0 -0
  127. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_hook_call_conventions.py +0 -0
  128. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_layer_kinds.py +0 -0
  129. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_logit_transform.py +0 -0
  130. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_mappers.py +0 -0
  131. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_mlp_internals.py +0 -0
  132. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_model_expectations.py +0 -0
  133. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_moe.py +0 -0
  134. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_multimodal_arch.py +0 -0
  135. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_new_models_gpu.py +0 -0
  136. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_no_chat_template.py +0 -0
  137. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_normalized_hook.py +0 -0
  138. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_notebook_stdout.py +0 -0
  139. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_packaging.py +0 -0
  140. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_parity_gpt2.py +0 -0
  141. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_per_layer_attn_dims.py +0 -0
  142. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_points_registry.py +0 -0
  143. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_protocol.py +0 -0
  144. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_published_benchmarks.py +0 -0
  145. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_qk_norm.py +0 -0
  146. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_qkv_layout.py +0 -0
  147. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_reasoning_spans.py +0 -0
  148. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_release.py +0 -0
  149. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_resid_mid.py +0 -0
  150. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_residual_basis.py +0 -0
  151. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_sandwich_norms.py +0 -0
  152. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_select.py +0 -0
  153. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_sliding_window_attn.py +0 -0
  154. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_small_models_gpu.py +0 -0
  155. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_static_dsv4_gpu.py +0 -0
  156. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_static_parity_gpu.py +0 -0
  157. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_static_warmup.py +0 -0
  158. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_steer_context.py +0 -0
  159. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_steer_math_parity.py +0 -0
  160. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_sync_loop.py +0 -0
  161. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_sync_parity.py +0 -0
  162. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_unified_free_functions.py +0 -0
  163. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_unresolved_families.py +0 -0
  164. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_capture_gpu.py +0 -0
  165. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_capture_scales.py +0 -0
  166. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_engine_loop.py +0 -0
  167. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_graph_path.py +0 -0
  168. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_graphs_on_gpu.py +0 -0
  169. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_hook_availability.py +0 -0
  170. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_hyper_connections.py +0 -0
  171. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_kv_isolation.py +0 -0
  172. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_new_points.py +0 -0
  173. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_only_families.py +0 -0
  174. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_plugin.py +0 -0
  175. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vllm_wire_grammar.py +0 -0
  176. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_vocabulary_boundary.py +0 -0
  177. {interp_engine-1.7.2 → interp_engine-1.8.0}/tests/test_worker_lens_capture_readout.py +0 -0
  178. {interp_engine-1.7.2 → interp_engine-1.8.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.2
3
+ Version: 1.8.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
@@ -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.
@@ -215,7 +225,9 @@ one entry in `chat_formatters.CODE_CHAT_FORMATS`.
215
225
 
216
226
  To attribute *tokens* to messages there are two methods, and the difference matters. `message_spans`
217
227
  gives per-token role, channel and section (`header` / `content` / `footer`), leaving the trailing
218
- generation scaffold owned by no message — use it to read or display structure. `message_partition`
228
+ generation scaffold owned by no message — use it to read or display structure. A system turn the
229
+ template injects on its own (Llama's knowledge-cutoff preamble, Qwen2.5's default persona) is
230
+ tagged `role="system"` with `message_index=None`, so it can be shown as its own turn. `message_partition`
219
231
  gives one contiguous `[start, end)` span per message that together cover every token, which is what
220
232
  mean-pooling activations per turn needs:
221
233
 
@@ -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,7 +136,21 @@ 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
155
  eager accelerate ``device_map="auto"``. Note that vLLM with ``num_gpus > 1``
92
156
  cannot serve per-head ``z`` or DFA, because attention heads are sharded across
@@ -118,7 +182,8 @@ def load_model(
118
182
  ValueError: ``backend`` is not one of :data:`BACKENDS`; or ``static_points`` /
119
183
  ``static_writes`` was passed on a backend other than ``"vllm-static"``; or
120
184
  ``backend="vllm-static"`` declared no taps at all; or ``enforce_eager=True`` was
121
- passed alongside a graph-replaying backend.
185
+ passed alongside a graph-replaying backend; or ``quantization`` / ``kv_cache_dtype``
186
+ asks the chosen backend for something it cannot apply.
122
187
  RuntimeError: a vLLM backend was requested but vLLM is not installed.
123
188
  GradientsUnsupported: ``requires_grad=True`` on a vLLM backend, which cannot
124
189
  provide gradients through its forward on any configuration.
@@ -182,6 +247,8 @@ def load_model(
182
247
  # installs nothing in Worker.load_model, and leaves hooks_available False.
183
248
  static_points, static_writes = [], None
184
249
 
250
+ _apply_load_precision(resolved, quantization, kv_cache_dtype, backend_kwargs)
251
+
185
252
  if use_vllm:
186
253
  require_vllm(f"backend={resolved!r} requested for {hf_model_id}")
187
254
  # `requires_grad` is an eager-only constructor kwarg, so on vLLM it would otherwise land as
@@ -44,6 +44,8 @@ from collections.abc import Iterator, Sequence
44
44
  from dataclasses import dataclass, field
45
45
  from typing import Any
46
46
 
47
+ from interp_engine.facts import FP16_EAGER_OVERFLOW_ARCHS
48
+
47
49
  GIB = 1024**3
48
50
 
49
51
  #: The backends :func:`estimate` understands. Mirrors ``load.BACKENDS`` deliberately rather than
@@ -145,6 +147,22 @@ CALIBRATION: dict[str, Calibration] = {
145
147
  "smallest of anything here."
146
148
  ),
147
149
  ),
150
+ "quant_on_load_gib": Calibration(
151
+ value=0.3,
152
+ unit="GiB",
153
+ why=(
154
+ "What vLLM charges against its budget for quantizing at load, past the narrowed tensors "
155
+ "themselves. It appears in vLLM's own 'weights' figure, so it comes out of the KV cache's "
156
+ "share the way the CUDA context does, and it does not scale with the model."
157
+ ),
158
+ source=(
159
+ "Qwen3-8B and Qwen3-4B, vllm, bf16 checkpoint with quantization='fp8', RTX 5090, "
160
+ "utilization 0.9. The fp8 tensors on the card summed to 8.80 and 4.12 GiB; vLLM reported "
161
+ "'Model loading took' 9.09 and 4.41. The same two models at bf16 reported their tensors "
162
+ "exactly, and the KV cache came out 0.98x of predicted on fp8 against 1.00x on bf16 -- the "
163
+ "0.30 GiB is the whole of that gap."
164
+ ),
165
+ ),
148
166
  "frag_fraction": Calibration(
149
167
  value=0.04,
150
168
  unit="fraction of card",
@@ -660,7 +678,7 @@ _SCHEME_WIDTH: tuple[tuple[tuple[str, ...], float], ...] = (
660
678
  )
661
679
 
662
680
 
663
- def _scheme_width(quant_method: str) -> float | None:
681
+ def scheme_width(quant_method: str) -> float | None:
664
682
  """Bytes per logical parameter for a quantization scheme, or None when it is not recognized.
665
683
 
666
684
  ``awq`` and ``gptq`` are listed at 4-bit because that is what they are in practice; an 8-bit GPTQ
@@ -676,6 +694,143 @@ def _scheme_width(quant_method: str) -> float | None:
676
694
  return None
677
695
 
678
696
 
697
+ # ------------------------------------------------------------ on-load quantization
698
+
699
+
700
+ @dataclass(frozen=True)
701
+ class OnLoadQuantization:
702
+ """A scheme a backend can apply to a wider checkpoint **at load**, with no calibration step.
703
+
704
+ This is what separates the rows below from AWQ, GPTQ, NVFP4 and MXFP4: those need a calibration
705
+ run and arrive as a repo, which :class:`WeightBytes` already prices from its headers. A row here is
706
+ one argument to ``load_model``, so a sizer can offer it on any bf16 checkpoint and the snippet it
707
+ prints still runs.
708
+ """
709
+
710
+ name: str
711
+ #: Bytes per **linear** weight once quantized, scales included. Embeddings and the unembed are not
712
+ #: linear layers to either quantizer and stay at the model dtype -- see :func:`on_load_weight_bytes`.
713
+ width: float
714
+ #: The backends that apply it. A backend missing here is refused, with the reason in :attr:`refusals`.
715
+ backends: tuple[str, ...]
716
+ #: What vLLM calls it in ``quantization=``; empty when no vLLM backend applies it.
717
+ vllm_name: str
718
+ #: The ``BitsAndBytesConfig`` arguments transformers needs; empty when eager does not apply it.
719
+ eager_config: dict[str, Any]
720
+ #: Why each absent backend cannot, and what to do instead. ``"*"`` covers every backend not named.
721
+ refusals: dict[str, str]
722
+ why: str
723
+
724
+ def applies_to(self, backend: str) -> bool:
725
+ return backend in self.backends
726
+
727
+ def refusal(self, backend: str) -> str:
728
+ """Why this backend cannot apply the scheme, or an empty string when it can."""
729
+ if self.applies_to(backend):
730
+ return ""
731
+ return self.refusals.get(backend) or self.refusals.get("*", f"{self.name} is not available on {backend}")
732
+
733
+
734
+ #: Every on-load scheme ``load_model(quantization=...)`` accepts and the sizer prices, by that name.
735
+ #:
736
+ #: The widths carry their scales. bitsandbytes' NF4 stores one fp32 absmax per 64-element block, which
737
+ #: is 0.0625 bytes a weight on top of the half byte; vLLM's fp8 and LLM.int8 keep one scale per output
738
+ #: channel, which rounds to nothing. FP8 on the eager backend is refused rather than priced:
739
+ #: transformers' ``FineGrainedFP8Config`` needs compute 8.9 and reaches for a DeepGEMM Hub kernel, so a
740
+ #: fit there would be a fit for two failure modes at once.
741
+ QUANTIZATIONS: dict[str, OnLoadQuantization] = {
742
+ "fp8": OnLoadQuantization(
743
+ name="fp8",
744
+ width=1.0,
745
+ backends=VLLM_BACKENDS,
746
+ vllm_name="fp8",
747
+ eager_config={},
748
+ refusals={
749
+ "eager": (
750
+ "transformers' on-load FP8 (FineGrainedFP8Config) needs compute 8.9 and a DeepGEMM Hub "
751
+ "kernel; load a checkpoint that ships in FP8 with dtype='auto' instead, or use a vLLM backend"
752
+ ),
753
+ },
754
+ why="dynamic per-channel FP8 on every linear layer: half of bf16, no calibration",
755
+ ),
756
+ "bnb-8bit": OnLoadQuantization(
757
+ name="bnb-8bit",
758
+ width=1.0,
759
+ backends=("eager",),
760
+ vllm_name="",
761
+ eager_config={"load_in_8bit": True},
762
+ refusals={
763
+ "*": (
764
+ "vLLM's in-flight bitsandbytes is 4-bit only; use quantization='bnb-4bit' or quantization='fp8' there"
765
+ ),
766
+ },
767
+ why="LLM.int8 through bitsandbytes: the one 8-bit scheme with a backward pass",
768
+ ),
769
+ "bnb-4bit": OnLoadQuantization(
770
+ name="bnb-4bit",
771
+ width=0.5625,
772
+ backends=(*VLLM_BACKENDS, "eager"),
773
+ vllm_name="bitsandbytes",
774
+ eager_config={"load_in_4bit": True, "bnb_4bit_quant_type": "nf4"},
775
+ refusals={},
776
+ why="NF4 through bitsandbytes, on both backends, and with a backward pass on eager",
777
+ ),
778
+ }
779
+
780
+
781
+ def quantization_refusal(quantization: str, backend: str) -> str:
782
+ """Why ``quantization`` cannot be applied on ``backend``, or an empty string when it can.
783
+
784
+ An unknown name is refused too, naming the ones that exist, so a typo does not price as "as stored".
785
+ """
786
+ if not quantization:
787
+ return ""
788
+ scheme = QUANTIZATIONS.get(quantization)
789
+ if scheme is None:
790
+ return f"unknown quantization {quantization!r}; the sizer prices {', '.join(QUANTIZATIONS)}"
791
+ return scheme.refusal(backend)
792
+
793
+
794
+ #: vLLM's spelling of each on-load scheme back to this module's, for the engine's own sizing in
795
+ #: ``vllm_capture/static.py``: ``quantization="bitsandbytes"`` there is in-flight NF4, ``bnb-4bit`` here.
796
+ VLLM_QUANTIZATION_NAMES: dict[str, str] = {
797
+ scheme.vllm_name: name for name, scheme in QUANTIZATIONS.items() if scheme.vllm_name
798
+ }
799
+
800
+
801
+ def narrow_linear_weights(base: int, *, narrow: float, stored: float, embedding_params: int) -> int:
802
+ """``base`` bytes of weights once every linear layer is narrowed from ``stored`` to ``narrow``.
803
+
804
+ **Neither vLLM nor bitsandbytes quantizes the embeddings or the unembed**, so those stay at the
805
+ stored width: on Llama-3.3-70B that is 2 x 1.05B parameters, about 3.9 GiB that "halve the
806
+ weights" would have counted away, and 6% of the answer. ``RedHatAI``'s FP8 export of the same
807
+ model shows the same shape from disk -- 67.68 GiB against 65.7 for 70.56B at one byte.
808
+
809
+ A scheme no narrower than the checkpoint changes nothing, which is the direction that OOMs.
810
+ """
811
+ if not base or narrow >= stored:
812
+ return int(base)
813
+ embeddings = min(max(int(embedding_params * stored), 0), int(base))
814
+ return int((int(base) - embeddings) * (narrow / stored) + embeddings)
815
+
816
+
817
+ def on_load_weight_bytes(facts: ModelMemoryFacts, dtype: str, quantization: str, *, dequantizes: bool) -> int:
818
+ """Weight bytes on the device once ``quantization`` has narrowed the checkpoint at load.
819
+
820
+ Starts from :meth:`WeightBytes.bytes_for_load`, which is the stored answer, and hands it to
821
+ :func:`narrow_linear_weights`. A checkpoint that already ships quantized is returned as stored:
822
+ a quantizer does not narrow what is already narrower, and :func:`estimate` says so in a warning.
823
+ """
824
+ weights = facts.weights
825
+ base = weights.bytes_for_load(dtype, dequantizes=dequantizes)
826
+ scheme = QUANTIZATIONS.get(quantization) if quantization else None
827
+ if scheme is None or weights.is_quantized:
828
+ return base
829
+ stored = dtype_bytes(dtype if dtype not in ("auto", "", None) else weights.stored_dtype)
830
+ embedding_params = facts.vocab_size * facts.d_model * (1 if facts.tied_embeddings else 2)
831
+ return narrow_linear_weights(base, narrow=scheme.width, stored=stored, embedding_params=embedding_params)
832
+
833
+
679
834
  def logical_param_count(elements_by_dtype: dict[str, int], quant_method: str = "", expert_dtype: str = "") -> int:
680
835
  """Logical parameters, unpacking whatever the containers hold.
681
836
 
@@ -700,7 +855,7 @@ def logical_param_count(elements_by_dtype: dict[str, int], quant_method: str = "
700
855
  An unquantized checkpoint is unaffected: its buckets are all float, so containers and parameters
701
856
  are the same thing.
702
857
  """
703
- widths = [w for w in (_scheme_width(quant_method), _scheme_width(expert_dtype)) if w is not None]
858
+ widths = [w for w in (scheme_width(quant_method), scheme_width(expert_dtype)) if w is not None]
704
859
  native = min(widths) if widths else None
705
860
  total = 0
706
861
  for tag, count in elements_by_dtype.items():
@@ -805,7 +960,7 @@ class WeightBytes:
805
960
  if self.is_quantized:
806
961
  if not dequantizes:
807
962
  return self.on_disk_bytes
808
- native = _scheme_width(self.quant_method)
963
+ native = scheme_width(self.quant_method)
809
964
  if native is not None and wanted <= native:
810
965
  # Asking for the width it is already stored at, or narrower than transformers will
811
966
  # give you: the checkpoint is served natively and the file size stands.
@@ -1015,7 +1170,7 @@ def _scheme_from_headers(elements_by_dtype: dict[str, int]) -> str:
1015
1170
  Requires the container buckets to hold a **majority** of the elements, so a small integer
1016
1171
  side-table cannot make a dense model look packed. Distinguishing 4-bit from 8-bit is then the
1017
1172
  presence of fp8 scales beside a byte payload, which is what every 4-bit export on the Hub looks
1018
- like; both answers are labels for :func:`_scheme_width`, not claims about a specific vendor
1173
+ like; both answers are labels for :func:`scheme_width`, not claims about a specific vendor
1019
1174
  format, so ``"nvfp4"`` here means "packed two-to-a-byte" rather than NVIDIA's exact encoding.
1020
1175
  """
1021
1176
  total = sum(elements_by_dtype.values())
@@ -1267,6 +1422,10 @@ class ModelMemoryFacts:
1267
1422
  #: NVIDIA's FP4 exports. Not the caller's ``kv_cache_dtype``: this is a property of the weights on
1268
1423
  #: disk, and vLLM honours it whether or not anyone asked. See :func:`hub_kv_quant_algo`.
1269
1424
  kv_quant_algo: str = ""
1425
+ #: Whether the unembed shares the embedding matrix. Read by :func:`on_load_weight_bytes`, since
1426
+ #: both stay at the model dtype under every on-load scheme and a tied pair is one matrix, not two.
1427
+ #: False when unknown, which charges the second matrix -- the direction that does not OOM.
1428
+ tied_embeddings: bool = False
1270
1429
 
1271
1430
  @property
1272
1431
  def kv_width(self) -> int:
@@ -1623,6 +1782,10 @@ class WorkloadSpec:
1623
1782
  #: default is ``"float32"``, not ``"auto"``, which doubles a bf16 checkpoint -- a sizer should
1624
1783
  #: say so rather than reproduce it silently.
1625
1784
  dtype: str = "auto"
1785
+ #: An on-load scheme from :data:`QUANTIZATIONS` (``"fp8"``, ``"bnb-4bit"``, ...), or empty for
1786
+ #: as stored. Not a dtype: vLLM rejects ``dtype="fp8"`` and transformers would store fp8 with no
1787
+ #: kernel behind it, which is why this is its own field and its own ``load_model`` argument.
1788
+ quantization: str = ""
1626
1789
  #: KV cache dtype. ``"auto"`` follows the model dtype.
1627
1790
  kv_cache_dtype: str = "auto"
1628
1791
  max_model_len: int = 0
@@ -2054,13 +2217,42 @@ def estimate(
2054
2217
  )
2055
2218
 
2056
2219
  # Only the eager backend expands a quantized checkpoint to the requested dtype; vLLM reads the same
2057
- # argument as an activation dtype and serves the packed weights. See `bytes_for_load`.
2058
- weights_total = facts.weights.bytes_for_load(spec.dtype, dequantizes=spec.backend == "eager")
2220
+ # argument as an activation dtype and serves the packed weights. See `bytes_for_load`. An on-load
2221
+ # scheme then narrows the linear layers and nothing else -- see `on_load_weight_bytes`.
2222
+ refused = quantization_refusal(spec.quantization, spec.backend)
2223
+ weights_total = on_load_weight_bytes(
2224
+ facts, spec.dtype, "" if refused else spec.quantization, dequantizes=spec.backend == "eager"
2225
+ )
2226
+ # True when the scheme took effect: the checkpoint was wider than it, and the backend runs it.
2227
+ narrowed = weights_total < facts.weights.bytes_for_load(spec.dtype, dequantizes=spec.backend == "eager")
2059
2228
  if not weights_total:
2060
2229
  warnings.append(
2061
2230
  "weight bytes are unknown, so every figure below is only the non-weight terms; "
2062
2231
  "pass hf_model_id or a config to model_memory_facts()"
2063
2232
  )
2233
+ if refused:
2234
+ warnings.append(f"quantization={spec.quantization!r} is refused on backend={spec.backend!r}: {refused}")
2235
+ elif spec.quantization and facts.weights.is_quantized:
2236
+ warnings.append(
2237
+ f"quantization={spec.quantization!r} changes nothing: this checkpoint already ships as "
2238
+ f"{facts.weights.quant_method}, and a quantizer does not narrow what is already narrower. "
2239
+ f"Priced as stored"
2240
+ )
2241
+ elif spec.quantization == "fp8" and not gpu.supports_fp8:
2242
+ warnings.append(
2243
+ f"{gpu.name} has no FP8 tensor cores, so quantization='fp8' runs weight-only through the "
2244
+ f"Marlin kernel: the same memory, but slower than on Ada/Hopper or newer"
2245
+ )
2246
+ if (
2247
+ spec.backend == "eager"
2248
+ and spec.dtype in ("float16", "fp16", "half")
2249
+ and (spec.attn_implementation or "eager") == "eager"
2250
+ and facts.architecture in FP16_EAGER_OVERFLOW_ARCHS
2251
+ ):
2252
+ warnings.append(
2253
+ f"{facts.architecture} overflows to NaN in float16 under attn_implementation='eager'; "
2254
+ f"use dtype='bfloat16' or attn_implementation='sdpa'"
2255
+ )
2064
2256
  if facts.weights.is_quantized:
2065
2257
  dequantized = facts.weights.dequantized_bytes() or 0
2066
2258
  if not gpu.supports_mxfp4_kernels and "fp4" in facts.weights.quant_method:
@@ -2103,7 +2295,7 @@ def estimate(
2103
2295
  "weights",
2104
2296
  per_card_weights,
2105
2297
  "eager",
2106
- f"{facts.weights.param_count / 1e9:.1f}B params at {spec.dtype}"
2298
+ _weights_note(spec, facts, refused)
2107
2299
  + (f", spread over {tp} GPUs" if tp > 1 else "")
2108
2300
  + f" [{facts.weights.source}]",
2109
2301
  )
@@ -2126,8 +2318,9 @@ def estimate(
2126
2318
 
2127
2319
  total = sum(term.bytes for term in terms)
2128
2320
  headroom = gpu.total_bytes - total
2129
- fits = headroom >= 0 and facts.trunk_dims_known
2130
- if not fits and facts.trunk_dims_known:
2321
+ # A refused quantization is a configuration that cannot start, whatever the arithmetic says.
2322
+ fits = headroom >= 0 and facts.trunk_dims_known and not refused
2323
+ if not fits and facts.trunk_dims_known and not refused:
2131
2324
  advice.extend(_eager_advice(activation, spec, facts))
2132
2325
  return MemoryEstimate(
2133
2326
  spec=spec,
@@ -2197,6 +2390,16 @@ def estimate(
2197
2390
  # Inside the pool. The CUDA context is here rather than outside because vLLM's budget is measured
2198
2391
  # against what the process is ALREADY using -- see CALIBRATION["cuda_context_gib"].
2199
2392
  terms.append(MemoryTerm("cuda_context", context, "pool", "process CUDA context, charged against vLLM's budget"))
2393
+ quant_charge = int(_cal("quant_on_load_gib") * GIB) if narrowed else 0
2394
+ if quant_charge:
2395
+ terms.append(
2396
+ MemoryTerm(
2397
+ "quant_on_load",
2398
+ quant_charge,
2399
+ "pool",
2400
+ f"what vLLM holds for quantizing to {spec.quantization} at load, past the tensors",
2401
+ )
2402
+ )
2200
2403
  if reserved_inside:
2201
2404
  terms.append(
2202
2405
  MemoryTerm(
@@ -2213,7 +2416,7 @@ def estimate(
2213
2416
  "weights",
2214
2417
  per_card_weights,
2215
2418
  "pool",
2216
- f"{facts.weights.param_count / 1e9:.1f}B params at {spec.dtype}"
2419
+ _weights_note(spec, facts, refused)
2217
2420
  + (f", sharded over TP={tp}" if tp > 1 else "")
2218
2421
  + f" [{facts.weights.source}]",
2219
2422
  )
@@ -2263,7 +2466,7 @@ def estimate(
2263
2466
 
2264
2467
  pool_available = int(spec.gpu_memory_utilization * gpu.total_bytes)
2265
2468
  outside_needed = overshoot + frag + reserved_outside
2266
- pool_needed = context + reserved_inside + per_card_weights + buffers + graphs + kv_floor
2469
+ pool_needed = context + quant_charge + reserved_inside + per_card_weights + buffers + graphs + kv_floor
2267
2470
 
2268
2471
  # Both constraints have to hold, and they fail differently -- see the module docstring.
2269
2472
  pool_headroom = pool_available - pool_needed
@@ -2275,15 +2478,21 @@ def estimate(
2275
2478
  # here; the warning above says which it is. An all-recurrent trunk reaches the same zero by a
2276
2479
  # different road -- there really is no cache -- and is refused for the same reason: what it holds
2277
2480
  # instead is a state pool nothing here prices.
2278
- fits = pool_headroom >= 0 and outside_headroom >= 0 and facts.trunk_dims_known and bool(facts.kv_caching_layers)
2481
+ fits = (
2482
+ pool_headroom >= 0
2483
+ and outside_headroom >= 0
2484
+ and facts.trunk_dims_known
2485
+ and bool(facts.kv_caching_layers)
2486
+ and not refused
2487
+ )
2279
2488
 
2280
- kv_room = max(pool_available - context - reserved_inside - per_card_weights - buffers - graphs, 0)
2489
+ kv_room = max(pool_available - context - quant_charge - reserved_inside - per_card_weights - buffers - graphs, 0)
2281
2490
  per_token = (
2282
2491
  kv_bytes_for_context(facts, spec.max_model_len, kv_dtype=spec.kv_cache_dtype, model_dtype=spec.dtype) / shards
2283
2492
  )
2284
2493
  kv_capacity = int(kv_room / (per_token / spec.max_model_len)) if per_token > 0 else 0
2285
2494
 
2286
- if not fits:
2495
+ if not fits and not refused:
2287
2496
  advice.extend(
2288
2497
  _vllm_advice(
2289
2498
  spec=spec,
@@ -2315,6 +2524,15 @@ def estimate(
2315
2524
  )
2316
2525
 
2317
2526
 
2527
+ def _weights_note(spec: WorkloadSpec, facts: ModelMemoryFacts, refused: str) -> str:
2528
+ """The first clause of the weights term: what is loaded, at what width."""
2529
+ note = f"{facts.weights.param_count / 1e9:.1f}B params at {spec.dtype}"
2530
+ if spec.quantization and not refused and not facts.weights.is_quantized:
2531
+ stored = spec.dtype if spec.dtype != "auto" else (facts.weights.stored_dtype or "bfloat16")
2532
+ note += f", {spec.quantization} on load (embeddings stay {stored})"
2533
+ return note
2534
+
2535
+
2318
2536
  def _eager_note(name: str, spec: WorkloadSpec, facts: ModelMemoryFacts) -> str:
2319
2537
  if name == "logits":
2320
2538
  return f"{spec.batch_size}x{spec.seq_len} x vocab {facts.vocab_size:,}, plus the fp32 upcast"
@@ -2394,9 +2612,14 @@ def _vllm_advice(
2394
2612
  out.append("dtype='bfloat16' halves the weights, which are the largest term in the pool")
2395
2613
  if weights > gpu.total_bytes * 0.7:
2396
2614
  need = weights / (gpu.total_bytes * 0.6)
2615
+ narrower = (
2616
+ "quantization='fp8' halves the linear layers on load"
2617
+ if not facts.weights.is_quantized and not spec.quantization
2618
+ else "use a quantized checkpoint"
2619
+ )
2397
2620
  out.append(
2398
2621
  f"the weights alone are {weights / GIB:.1f} GiB of a {gpu.total_gib:.1f} GiB card; "
2399
- f"num_gpus={max(2, int(need) + 1)} shards them, or use a quantized checkpoint"
2622
+ f"num_gpus={max(2, int(need) + 1)} shards them, or {narrower}"
2400
2623
  )
2401
2624
  if spec.graphs_on:
2402
2625
  out.append(
@@ -2477,6 +2700,7 @@ def model_memory_facts(
2477
2700
  max_position_embeddings=int(getattr(cfg, "max_position_embeddings", 0) or 0),
2478
2701
  architecture=f.architecture,
2479
2702
  kv_quant_algo=hub_kv_quant_algo(hf_model_id, token),
2703
+ tied_embeddings=f.tied_embeddings,
2480
2704
  )
2481
2705
 
2482
2706
 
@@ -2507,6 +2731,7 @@ def fit(
2507
2731
  reservations: Reservations | None = None,
2508
2732
  max_model_len: int = 0,
2509
2733
  dtype: str = "auto",
2734
+ quantization: str = "",
2510
2735
  kv_cache_dtype: str = "auto",
2511
2736
  num_gpus: int = 1,
2512
2737
  static_sites: int = 0,
@@ -2521,6 +2746,9 @@ def fit(
2521
2746
  ) -> tuple[WorkloadSpec, MemoryEstimate] | None:
2522
2747
  """The largest configuration of this shape that fits, or None when none does.
2523
2748
 
2749
+ A ``quantization`` the backend cannot apply returns None before any ladder is walked: the estimate
2750
+ would price the refusal as a warning and refuse to fit, so every rung would come back False.
2751
+
2524
2752
  Solves in the order the constraints actually bind:
2525
2753
 
2526
2754
  1. **Utilization.** Everything outside the pool is fixed by the card and the reservations, so the
@@ -2537,7 +2765,7 @@ def fit(
2537
2765
  """
2538
2766
  res = reservations or Reservations()
2539
2767
 
2540
- if not facts.trunk_dims_known:
2768
+ if not facts.trunk_dims_known or quantization_refusal(quantization, backend):
2541
2769
  # Short-circuit what would happen anyway. `estimate` refuses to report `fits` on a model whose
2542
2770
  # dims are unknown, so every rung of every ladder below would come back False; returning here
2543
2771
  # just saves walking them. Callers wanting to tell "cannot size" from "does not fit" should
@@ -2566,6 +2794,7 @@ def fit(
2566
2794
  spec = WorkloadSpec(
2567
2795
  backend="eager",
2568
2796
  dtype=dtype,
2797
+ quantization=quantization,
2569
2798
  num_gpus=num_gpus,
2570
2799
  batch_size=batch_size,
2571
2800
  seq_len=prompt,
@@ -2611,6 +2840,7 @@ def fit(
2611
2840
  spec = WorkloadSpec(
2612
2841
  backend=backend,
2613
2842
  dtype=dtype,
2843
+ quantization=quantization,
2614
2844
  kv_cache_dtype=kv_cache_dtype,
2615
2845
  max_model_len=context,
2616
2846
  # A prefill batch wider than the context is waste: nothing can fill it.