infer-stack 0.6.0__py3-none-any.whl

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 (44) hide show
  1. infer_stack/__init__.py +2 -0
  2. infer_stack/backends/__init__.py +7 -0
  3. infer_stack/backends/compose_renderer.py +243 -0
  4. infer_stack/backends/kubeai_renderer.py +202 -0
  5. infer_stack/benchmark.py +38 -0
  6. infer_stack/catalog.py +438 -0
  7. infer_stack/cli/__init__.py +169 -0
  8. infer_stack/cli/__main__.py +4 -0
  9. infer_stack/cli/commands_profile.py +467 -0
  10. infer_stack/cli/commands_runtime.py +719 -0
  11. infer_stack/cli/commands_smoke.py +691 -0
  12. infer_stack/cli/compose.py +755 -0
  13. infer_stack/cli/context.py +471 -0
  14. infer_stack/cli/options.py +134 -0
  15. infer_stack/cli/probes.py +178 -0
  16. infer_stack/config.py +450 -0
  17. infer_stack/contracts.py +223 -0
  18. infer_stack/diff_prompt.py +117 -0
  19. infer_stack/docker_utils.py +230 -0
  20. infer_stack/env_utils.py +97 -0
  21. infer_stack/experimental/model_catalog_discover.py +1155 -0
  22. infer_stack/experimental/model_memory_estimator.py +1264 -0
  23. infer_stack/experimental/stress_test_long_context.py +397 -0
  24. infer_stack/hardware.py +70 -0
  25. infer_stack/kubeai_ops.py +76 -0
  26. infer_stack/paths.py +87 -0
  27. infer_stack/profile_runtime.py +46 -0
  28. infer_stack/renderer.py +19 -0
  29. infer_stack/resolver.py +1092 -0
  30. infer_stack/templates/default-models.yaml +674 -0
  31. infer_stack/templates/default-ollama-models.yaml +31 -0
  32. infer_stack/templates/default-profiles.yaml +1731 -0
  33. infer_stack/templates/default-vllm-models.yaml +714 -0
  34. infer_stack/templates/docker-compose.yml.j2 +430 -0
  35. infer_stack/templates/litellm_config.yaml.j2 +44 -0
  36. infer_stack/templates/nginx.conf.j2 +84 -0
  37. infer_stack/tuning.py +3 -0
  38. infer_stack/validator.py +314 -0
  39. infer_stack/verification.py +46 -0
  40. infer_stack-0.6.0.dist-info/METADATA +1034 -0
  41. infer_stack-0.6.0.dist-info/RECORD +44 -0
  42. infer_stack-0.6.0.dist-info/WHEEL +5 -0
  43. infer_stack-0.6.0.dist-info/entry_points.txt +2 -0
  44. infer_stack-0.6.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,1264 @@
1
+ from __future__ import annotations
2
+
3
+ import argparse
4
+ import json
5
+ import math
6
+ import re
7
+ from dataclasses import dataclass, field
8
+ from pathlib import Path
9
+ from typing import Any, Final, Literal, Pattern, Sequence
10
+
11
+ from rich.console import Console
12
+ from rich.table import Table
13
+ from rich.panel import Panel
14
+ from rich import box
15
+
16
+ try:
17
+ from huggingface_hub import HfApi, hf_hub_download
18
+ except Exception: # pragma: no cover
19
+ HfApi = None # type: ignore
20
+ hf_hub_download = None # type: ignore
21
+
22
+
23
+ GiB = 1024**3
24
+
25
+
26
+ # --- model presets ---------------------------------------------------------
27
+ #
28
+ # Keep the presets we've been interested in, but use if 0 / if 1 gates so it
29
+ # is obvious which ones are active by default right now.
30
+
31
+ DEFAULT_PRESET_REPO_IDS: list[str] = []
32
+
33
+ if 1:
34
+ DEFAULT_PRESET_REPO_IDS += [
35
+ 'Qwen/Qwen3.5-122B-A10B',
36
+ 'Qwen/Qwen3.5-122B-A10B-FP8',
37
+ # Official integer-quantized preset currently available on HF.
38
+ 'Qwen/Qwen3.5-122B-A10B-GPTQ-Int4',
39
+ ]
40
+
41
+ if 0:
42
+ DEFAULT_PRESET_REPO_IDS += [
43
+ 'Qwen/Qwen3.5-2B',
44
+ 'Qwen/Qwen3.5-9B',
45
+ 'Qwen/Qwen3.5-27B',
46
+ 'Qwen/Qwen3.5-35B-A3B',
47
+ ]
48
+
49
+ if 0:
50
+ DEFAULT_PRESET_REPO_IDS += [
51
+ 'Qwen/Qwen3.6-30B-A3B-Thinking-2507',
52
+ 'Qwen/Qwen3.6-235B-A22B-Thinking-2507',
53
+ 'google/gemma-4-27b-it',
54
+ 'google/gemma-4-9b-it',
55
+ ]
56
+
57
+
58
+ @dataclass(frozen=True)
59
+ class ModelDiscoveryRule:
60
+ author: str
61
+ search: str
62
+ include: Pattern[str]
63
+
64
+
65
+ # Keep discovery rules around for optional no-hardcoded discovery workflows.
66
+ DEFAULT_MODEL_DISCOVERY_RULES: Final[tuple[ModelDiscoveryRule, ...]] = (
67
+ ModelDiscoveryRule(
68
+ author='Qwen',
69
+ search='Qwen3.5',
70
+ include=re.compile(r'^Qwen3\.5-.*$', re.IGNORECASE),
71
+ ),
72
+ ModelDiscoveryRule(
73
+ author='Qwen',
74
+ search='Qwen3.6',
75
+ include=re.compile(r'^Qwen3\.6-.*$', re.IGNORECASE),
76
+ ),
77
+ ModelDiscoveryRule(
78
+ author='google',
79
+ search='gemma-4',
80
+ include=re.compile(r'^gemma-4-.*$', re.IGNORECASE),
81
+ ),
82
+ )
83
+ DEFAULT_MODEL_DISCOVERY_LIMIT = 100
84
+
85
+
86
+ FIT_STYLES: Final[dict[str, str]] = {
87
+ 'yes': 'bold green',
88
+ 'no': 'bold red',
89
+ 'maybe': 'bold yellow',
90
+ }
91
+
92
+
93
+ def fit_status_markup(status: Literal['yes', 'no', 'maybe']) -> str:
94
+ style = FIT_STYLES[status]
95
+ return f'[{style}]{status}[/{style}]'
96
+
97
+
98
+ def gib(x: float) -> float:
99
+ return x / GiB
100
+
101
+
102
+ def gib_text(x: float | None) -> str:
103
+ if x is None:
104
+ return '-'
105
+ return f'{gib(x):.2f}'
106
+
107
+
108
+ def range_text(lo: float, hi: float) -> str:
109
+ return f'{gib(lo):.2f}–{gib(hi):.2f}'
110
+
111
+
112
+ @dataclass(frozen=True)
113
+ class DTypeSpec:
114
+ name: str
115
+ bytes_per_element: float
116
+ notes: str = ''
117
+
118
+
119
+ DTYPE_F16 = DTypeSpec('fp16/bf16', 2.0, 'Standard half precision')
120
+ DTYPE_FP8 = DTypeSpec('fp8', 1.0, 'Approximate FP8 storage')
121
+ DTYPE_F32 = DTypeSpec('fp32', 4.0, 'Single precision')
122
+
123
+
124
+ @dataclass(frozen=True)
125
+ class WeightFootprint:
126
+ total_bytes: int | None
127
+ source: str
128
+ notes: str = ''
129
+
130
+
131
+ @dataclass(frozen=True)
132
+ class CacheGroupSpec:
133
+ name: str
134
+ kind: Literal['full_kv', 'sliding_kv', 'linear_recurrent', 'linear_conv']
135
+ layer_type: str
136
+ total_layers: int
137
+ unique_cache_layers: int
138
+ num_heads: int | None = None
139
+ head_dim: int | None = None
140
+ seq_len_mode: Literal['full', 'sliding', 'fixed'] = 'full'
141
+ sliding_window: int | None = None
142
+ fixed_elements_per_sequence: int | None = None
143
+ kv_copies: int = 2
144
+ dtype_source: Literal['kv', 'linear_state'] = 'kv'
145
+ notes: str = ''
146
+
147
+ def request_floor_elements_total(
148
+ self, deployment: 'DeploymentSpec'
149
+ ) -> float:
150
+ batch = deployment.concurrent_sequences
151
+ if self.kind in {'full_kv', 'sliding_kv'}:
152
+ if self.num_heads is None or self.head_dim is None:
153
+ raise ValueError(
154
+ f'Missing num_heads/head_dim for cache group {self.name}'
155
+ )
156
+ seq_len = deployment.total_sequence_tokens
157
+ if self.seq_len_mode == 'sliding':
158
+ if self.sliding_window is None:
159
+ raise ValueError(f'Sliding window not set for {self.name}')
160
+ seq_len = min(seq_len, self.sliding_window)
161
+ return (
162
+ self.kv_copies
163
+ * self.unique_cache_layers
164
+ * batch
165
+ * seq_len
166
+ * self.num_heads
167
+ * self.head_dim
168
+ )
169
+ if self.fixed_elements_per_sequence is None:
170
+ raise ValueError(
171
+ f'Missing fixed_elements_per_sequence for cache group {self.name}'
172
+ )
173
+ return (
174
+ self.unique_cache_layers * batch * self.fixed_elements_per_sequence
175
+ )
176
+
177
+ def bytes_per_request_cluster(self, deployment: 'DeploymentSpec') -> float:
178
+ dtype = (
179
+ deployment.kv_cache_dtype
180
+ if self.dtype_source == 'kv'
181
+ else deployment.linear_state_dtype
182
+ )
183
+ return (
184
+ self.request_floor_elements_total(deployment)
185
+ * dtype.bytes_per_element
186
+ )
187
+
188
+ def single_sequence_token_slope_cluster(
189
+ self, deployment: 'DeploymentSpec'
190
+ ) -> float:
191
+ dtype = (
192
+ deployment.kv_cache_dtype
193
+ if self.dtype_source == 'kv'
194
+ else deployment.linear_state_dtype
195
+ )
196
+ if self.kind in {'full_kv', 'sliding_kv'}:
197
+ if self.num_heads is None or self.head_dim is None:
198
+ raise ValueError(
199
+ f'Missing num_heads/head_dim for cache group {self.name}'
200
+ )
201
+ return (
202
+ self.kv_copies
203
+ * self.unique_cache_layers
204
+ * self.num_heads
205
+ * self.head_dim
206
+ * dtype.bytes_per_element
207
+ )
208
+ return 0.0
209
+
210
+ def single_sequence_fixed_bytes_cluster(
211
+ self, deployment: 'DeploymentSpec'
212
+ ) -> float:
213
+ dtype = (
214
+ deployment.kv_cache_dtype
215
+ if self.dtype_source == 'kv'
216
+ else deployment.linear_state_dtype
217
+ )
218
+ if self.kind in {'full_kv', 'sliding_kv'}:
219
+ return 0.0
220
+ if self.fixed_elements_per_sequence is None:
221
+ raise ValueError(
222
+ f'Missing fixed_elements_per_sequence for cache group {self.name}'
223
+ )
224
+ return (
225
+ self.unique_cache_layers
226
+ * self.fixed_elements_per_sequence
227
+ * dtype.bytes_per_element
228
+ )
229
+
230
+
231
+ @dataclass(frozen=True)
232
+ class ModelMemorySpec:
233
+ repo_id: str
234
+ family: Literal['qwen3.5', 'qwen3.6', 'gemma4']
235
+ architecture: str
236
+ max_position_embeddings: int
237
+ text_hidden_size: int
238
+ num_hidden_layers: int
239
+ layer_types: tuple[str, ...]
240
+ weight_footprint: WeightFootprint
241
+ cache_groups: tuple[CacheGroupSpec, ...]
242
+ has_vision: bool = False
243
+ notes: tuple[str, ...] = field(default_factory=tuple)
244
+
245
+
246
+ @dataclass(frozen=True)
247
+ class StartupOverheadInterval:
248
+ low_bytes_per_gpu: float
249
+ high_bytes_per_gpu: float
250
+ notes: str = ''
251
+
252
+
253
+ @dataclass(frozen=True)
254
+ class StartupOverheadPolicy:
255
+ base_runtime_gib_low: float = 2.25
256
+ base_runtime_gib_high: float = 3.25
257
+ compile_graph_gib_low: float = 1.00
258
+ compile_graph_gib_high: float = 1.75
259
+ warmup_workspace_gib_low: float = 0.75
260
+ warmup_workspace_gib_high: float = 1.25
261
+ hidden_size_scale_low_gib_per_8k: float = 0.25
262
+ hidden_size_scale_high_gib_per_8k: float = 0.50
263
+ multimodal_margin_gib_low: float = 1.50
264
+ multimodal_margin_gib_high: float = 3.50
265
+
266
+ def estimate(
267
+ self, model: ModelMemorySpec, deployment: 'DeploymentSpec'
268
+ ) -> StartupOverheadInterval:
269
+ size_ratio = model.text_hidden_size / 8192.0
270
+ low = (
271
+ self.base_runtime_gib_low
272
+ + self.compile_graph_gib_low
273
+ + self.warmup_workspace_gib_low
274
+ + self.hidden_size_scale_low_gib_per_8k * size_ratio
275
+ ) * GiB
276
+ high = (
277
+ self.base_runtime_gib_high
278
+ + self.compile_graph_gib_high
279
+ + self.warmup_workspace_gib_high
280
+ + self.hidden_size_scale_high_gib_per_8k * size_ratio
281
+ ) * GiB
282
+ notes = [
283
+ 'Deterministic startup-overhead interval covering runtime resident allocations, compile/cudagraph setup, and warmup workspaces.'
284
+ ]
285
+ if model.has_vision and not deployment.language_model_only:
286
+ low += self.multimodal_margin_gib_low * GiB
287
+ high += self.multimodal_margin_gib_high * GiB
288
+ notes.append(
289
+ 'Includes multimodal resident-memory margin because language_model_only is disabled.'
290
+ )
291
+ elif model.has_vision and deployment.language_model_only:
292
+ notes.append(
293
+ 'Multimodal resident-memory margin omitted because language_model_only is enabled.'
294
+ )
295
+ return StartupOverheadInterval(low, high, ' '.join(notes))
296
+
297
+
298
+ DEFAULT_STARTUP_OVERHEAD_POLICY = StartupOverheadPolicy()
299
+
300
+
301
+ @dataclass(frozen=True)
302
+ class DeploymentSpec:
303
+ name: str
304
+ tensor_parallel_size: int
305
+ data_parallel_size: int = 1
306
+ concurrent_sequences: int = 1
307
+ prompt_text_tokens: int = 8192
308
+ media_soft_tokens: int = 0
309
+ max_new_tokens: int = 0
310
+ kv_cache_dtype: DTypeSpec = DTYPE_F16
311
+ linear_state_dtype: DTypeSpec = DTYPE_F16
312
+ gpu_memory_bytes: int | None = None
313
+ gpu_memory_utilization: float = 0.95
314
+ language_model_only: bool = True
315
+ startup_overhead_policy: StartupOverheadPolicy = (
316
+ DEFAULT_STARTUP_OVERHEAD_POLICY
317
+ )
318
+
319
+ @property
320
+ def total_sequence_tokens(self) -> int:
321
+ return (
322
+ self.prompt_text_tokens
323
+ + self.media_soft_tokens
324
+ + self.max_new_tokens
325
+ )
326
+
327
+ @property
328
+ def managed_budget_bytes_per_gpu(self) -> float | None:
329
+ if self.gpu_memory_bytes is None:
330
+ return None
331
+ return self.gpu_memory_bytes * self.gpu_memory_utilization
332
+
333
+ @property
334
+ def gpu_memory_gib(self) -> float | None:
335
+ if self.gpu_memory_bytes is None:
336
+ return None
337
+ return self.gpu_memory_bytes / GiB
338
+
339
+
340
+ @dataclass(frozen=True)
341
+ class RequestFloor:
342
+ total_bytes_per_gpu: float
343
+ token_slope_bytes_per_gpu: float
344
+ fixed_bytes_per_gpu: float
345
+ notes: tuple[str, ...] = field(default_factory=tuple)
346
+
347
+
348
+ @dataclass(frozen=True)
349
+ class StartupFit:
350
+ weight_bytes_per_gpu: float | None
351
+ overhead_low_bytes_per_gpu: float
352
+ overhead_high_bytes_per_gpu: float
353
+ used_low_bytes_per_gpu: float | None
354
+ used_high_bytes_per_gpu: float | None
355
+ margin_low_bytes_per_gpu: float | None
356
+ margin_high_bytes_per_gpu: float | None
357
+ status: Literal['yes', 'no', 'maybe']
358
+
359
+
360
+ @dataclass(frozen=True)
361
+ class CapacityEstimate:
362
+ kv_budget_low_bytes_per_gpu: float | None
363
+ kv_budget_high_bytes_per_gpu: float | None
364
+ kv_tokens_low_cluster: float | None
365
+ kv_tokens_high_cluster: float | None
366
+ max_concurrency_low: float | None
367
+ max_concurrency_high: float | None
368
+ fits_target_low: bool | None
369
+ fits_target_high: bool | None
370
+
371
+
372
+ @dataclass(frozen=True)
373
+ class MemoryEstimate:
374
+ model: ModelMemorySpec
375
+ deployment: DeploymentSpec
376
+ request_floor: RequestFloor
377
+ startup_fit: StartupFit
378
+ capacity: CapacityEstimate
379
+
380
+
381
+ def _load_json(
382
+ repo_id: str, filename: str, token: str | None = None
383
+ ) -> dict[str, Any]:
384
+ if hf_hub_download is None:
385
+ raise RuntimeError('huggingface_hub is required to fetch configs')
386
+ path = hf_hub_download(repo_id, filename=filename, token=token)
387
+ return json.loads(Path(path).read_text())
388
+
389
+
390
+ def _model_info(
391
+ repo_id: str, token: str | None = None
392
+ ) -> tuple[Any | None, Any | None]:
393
+ if HfApi is None:
394
+ return None, None
395
+ api = HfApi(token=token)
396
+ info_expand = None
397
+ info_files = None
398
+ try:
399
+ info_expand = api.model_info(
400
+ repo_id,
401
+ expand=[
402
+ 'config',
403
+ 'safetensors',
404
+ 'siblings',
405
+ 'tags',
406
+ 'pipeline_tag',
407
+ 'createdAt',
408
+ 'lastModified',
409
+ ],
410
+ )
411
+ except Exception:
412
+ try:
413
+ info_expand = api.model_info(repo_id)
414
+ except Exception:
415
+ info_expand = None
416
+ try:
417
+ info_files = api.model_info(repo_id, files_metadata=True)
418
+ except Exception:
419
+ info_files = None
420
+ return info_expand, info_files
421
+
422
+
423
+ def _to_dict(obj: Any) -> dict[str, Any]:
424
+ if obj is None:
425
+ return {}
426
+ if isinstance(obj, dict):
427
+ return obj
428
+ data = getattr(obj, '__dict__', None)
429
+ if isinstance(data, dict):
430
+ return dict(data)
431
+ out: dict[str, Any] = {}
432
+ for name in dir(obj):
433
+ if name.startswith('_'):
434
+ continue
435
+ value = getattr(obj, name)
436
+ if callable(value):
437
+ continue
438
+ out[name] = value
439
+ return out
440
+
441
+
442
+ def _listify(x: Any) -> list[Any]:
443
+ if x is None:
444
+ return []
445
+ if isinstance(x, list):
446
+ return x
447
+ if isinstance(x, tuple):
448
+ return list(x)
449
+ return [x]
450
+
451
+
452
+ def _weight_footprint_from_hf(
453
+ repo_id: str, token: str | None = None
454
+ ) -> WeightFootprint:
455
+ info_expand, info_files = _model_info(repo_id, token=token)
456
+ expand_dict = _to_dict(info_expand)
457
+ candidates: list[tuple[int, str, str]] = []
458
+
459
+ safetensors_info = expand_dict.get('safetensors')
460
+ if isinstance(safetensors_info, dict):
461
+ total = safetensors_info.get('total') or safetensors_info.get(
462
+ 'total_size'
463
+ )
464
+ if isinstance(total, (int, float)):
465
+ candidates.append(
466
+ (
467
+ int(total),
468
+ 'huggingface model_info.safetensors.total',
469
+ 'Checkpoint storage bytes reported by HF safetensors metadata.',
470
+ )
471
+ )
472
+
473
+ if info_files is not None:
474
+ siblings = _listify(getattr(info_files, 'siblings', None))
475
+ total = 0
476
+ used = False
477
+ for sibling in siblings:
478
+ sdict = _to_dict(sibling)
479
+ name = str(sdict.get('rfilename') or sdict.get('path') or '')
480
+ size = sdict.get('size')
481
+ if name.endswith(('.safetensors', '.bin', '.pt')) and isinstance(
482
+ size, (int, float)
483
+ ):
484
+ total += int(size)
485
+ used = True
486
+ if used:
487
+ candidates.append(
488
+ (
489
+ total,
490
+ 'sum(model_info(files_metadata=True).siblings[*].size)',
491
+ 'Checkpoint storage bytes reconstructed from HF file metadata.',
492
+ )
493
+ )
494
+
495
+ if not candidates:
496
+ return WeightFootprint(
497
+ total_bytes=None,
498
+ source='unavailable',
499
+ notes='Could not recover checkpoint byte size from HF metadata.',
500
+ )
501
+
502
+ best_total, best_source, best_note = max(candidates, key=lambda x: x[0])
503
+ notes = [best_note]
504
+ if len(candidates) > 1:
505
+ rendered = ', '.join(
506
+ f'{src_name}={total / GiB:.2f} GiB'
507
+ for total, src_name, _ in candidates
508
+ )
509
+ notes.append(
510
+ 'Multiple HF byte signals were available; using the maximum to avoid undercounting. '
511
+ f'Candidates: {rendered}.'
512
+ )
513
+
514
+ return WeightFootprint(
515
+ total_bytes=best_total, source=best_source, notes=' '.join(notes)
516
+ )
517
+
518
+
519
+ def _default_qwen_layer_types(
520
+ num_hidden_layers: int, full_attention_interval: int = 4
521
+ ) -> list[str]:
522
+ return [
523
+ 'linear_attention'
524
+ if bool((i + 1) % full_attention_interval)
525
+ else 'full_attention'
526
+ for i in range(num_hidden_layers)
527
+ ]
528
+
529
+
530
+ def _default_gemma_layer_types(
531
+ num_hidden_layers: int, sliding_pattern: int = 6
532
+ ) -> list[str]:
533
+ layer_types = [
534
+ 'sliding_attention'
535
+ if bool((i + 1) % sliding_pattern)
536
+ else 'full_attention'
537
+ for i in range(num_hidden_layers)
538
+ ]
539
+ if layer_types and layer_types[-1] != 'full_attention':
540
+ layer_types[-1] = 'full_attention'
541
+ return layer_types
542
+
543
+
544
+ def _build_qwen_spec(
545
+ repo_id: str,
546
+ raw_config: dict[str, Any],
547
+ weight_footprint: WeightFootprint,
548
+ family: Literal['qwen3.5', 'qwen3.6'],
549
+ ) -> ModelMemorySpec:
550
+ text = raw_config.get('text_config', raw_config)
551
+ vision = raw_config.get('vision_config', {})
552
+ num_hidden_layers = int(text['num_hidden_layers'])
553
+ layer_types = tuple(
554
+ text.get('layer_types') or _default_qwen_layer_types(num_hidden_layers)
555
+ )
556
+ num_full = sum(1 for t in layer_types if t == 'full_attention')
557
+ num_linear = sum(1 for t in layer_types if t == 'linear_attention')
558
+
559
+ full_group = CacheGroupSpec(
560
+ name='full_attention_kv',
561
+ kind='full_kv',
562
+ layer_type='full_attention',
563
+ total_layers=num_full,
564
+ unique_cache_layers=num_full,
565
+ num_heads=int(text['num_key_value_heads']),
566
+ head_dim=int(
567
+ text.get('head_dim')
568
+ or (text['hidden_size'] // text['num_attention_heads'])
569
+ ),
570
+ seq_len_mode='full',
571
+ kv_copies=2,
572
+ dtype_source='kv',
573
+ notes='Standard KV cache for the full-attention layers only.',
574
+ )
575
+
576
+ linear_num_key_heads = int(text['linear_num_key_heads'])
577
+ linear_num_value_heads = int(text['linear_num_value_heads'])
578
+ linear_key_head_dim = int(text['linear_key_head_dim'])
579
+ linear_value_head_dim = int(text['linear_value_head_dim'])
580
+ conv_kernel = int(text['linear_conv_kernel_dim'])
581
+ conv_width = (
582
+ 2 * linear_num_key_heads * linear_key_head_dim
583
+ + linear_num_value_heads * linear_value_head_dim
584
+ )
585
+
586
+ linear_recurrent = CacheGroupSpec(
587
+ name='linear_attention_recurrent_state',
588
+ kind='linear_recurrent',
589
+ layer_type='linear_attention',
590
+ total_layers=num_linear,
591
+ unique_cache_layers=num_linear,
592
+ fixed_elements_per_sequence=linear_num_value_heads
593
+ * linear_key_head_dim
594
+ * linear_value_head_dim,
595
+ seq_len_mode='fixed',
596
+ dtype_source='linear_state',
597
+ notes='Fixed-size recurrent state for Qwen Gated DeltaNet layers.',
598
+ )
599
+ linear_conv = CacheGroupSpec(
600
+ name='linear_attention_conv_state',
601
+ kind='linear_conv',
602
+ layer_type='linear_attention',
603
+ total_layers=num_linear,
604
+ unique_cache_layers=num_linear,
605
+ fixed_elements_per_sequence=conv_width * conv_kernel,
606
+ seq_len_mode='fixed',
607
+ dtype_source='linear_state',
608
+ notes='Fixed-size causal-convolution state for Qwen Gated DeltaNet layers.',
609
+ )
610
+
611
+ notes = [
612
+ f'Parsed as {family} hybrid text+vision model.',
613
+ f'Layer mix: {num_linear} linear-attention layers and {num_full} full-attention layers.',
614
+ ]
615
+
616
+ has_vision = isinstance(vision, dict) and bool(vision)
617
+
618
+ return ModelMemorySpec(
619
+ repo_id=repo_id,
620
+ family=family,
621
+ architecture='hybrid_qwen_linear_plus_full_attention',
622
+ max_position_embeddings=int(text['max_position_embeddings']),
623
+ text_hidden_size=int(text['hidden_size']),
624
+ num_hidden_layers=num_hidden_layers,
625
+ layer_types=layer_types,
626
+ weight_footprint=weight_footprint,
627
+ cache_groups=(full_group, linear_recurrent, linear_conv),
628
+ has_vision=has_vision,
629
+ notes=tuple(notes),
630
+ )
631
+
632
+
633
+ def _build_gemma4_spec(
634
+ repo_id: str, raw_config: dict[str, Any], weight_footprint: WeightFootprint
635
+ ) -> ModelMemorySpec:
636
+ text = raw_config.get('text_config', raw_config)
637
+ layer_types = tuple(
638
+ text.get('layer_types')
639
+ or _default_gemma_layer_types(int(text['num_hidden_layers']))
640
+ )
641
+ shared_tail = int(text.get('num_kv_shared_layers', 0))
642
+ shared_mask = [False] * len(layer_types)
643
+ for i in range(max(0, len(layer_types) - shared_tail), len(layer_types)):
644
+ shared_mask[i] = True
645
+
646
+ num_heads_sliding = int(text['num_key_value_heads'])
647
+ head_dim_sliding = int(
648
+ text.get('head_dim')
649
+ or (text['hidden_size'] // text['num_attention_heads'])
650
+ )
651
+ num_heads_full = int(
652
+ text.get('num_global_key_value_heads')
653
+ or text.get('num_key_value_heads')
654
+ )
655
+ head_dim_full = int(
656
+ text.get('global_head_dim')
657
+ or text.get('head_dim')
658
+ or (text['hidden_size'] // text['num_attention_heads'])
659
+ )
660
+ kv_copies = 1 if bool(text.get('attention_k_eq_v', False)) else 2
661
+
662
+ sliding_total = sum(1 for t in layer_types if t == 'sliding_attention')
663
+ full_total = sum(1 for t in layer_types if t == 'full_attention')
664
+ sliding_unique = sum(
665
+ 1
666
+ for i, t in enumerate(layer_types)
667
+ if t == 'sliding_attention' and not shared_mask[i]
668
+ )
669
+ full_unique = sum(
670
+ 1
671
+ for i, t in enumerate(layer_types)
672
+ if t == 'full_attention' and not shared_mask[i]
673
+ )
674
+
675
+ sliding_group = CacheGroupSpec(
676
+ name='sliding_attention_kv',
677
+ kind='sliding_kv',
678
+ layer_type='sliding_attention',
679
+ total_layers=sliding_total,
680
+ unique_cache_layers=sliding_unique,
681
+ num_heads=num_heads_sliding,
682
+ head_dim=head_dim_sliding,
683
+ seq_len_mode='sliding',
684
+ sliding_window=int(text['sliding_window']),
685
+ kv_copies=kv_copies,
686
+ dtype_source='kv',
687
+ notes='Sliding-window KV cache.',
688
+ )
689
+ full_group = CacheGroupSpec(
690
+ name='full_attention_kv',
691
+ kind='full_kv',
692
+ layer_type='full_attention',
693
+ total_layers=full_total,
694
+ unique_cache_layers=full_unique,
695
+ num_heads=num_heads_full,
696
+ head_dim=head_dim_full,
697
+ seq_len_mode='full',
698
+ kv_copies=kv_copies,
699
+ dtype_source='kv',
700
+ notes='Global/full-attention KV cache.',
701
+ )
702
+
703
+ vision = raw_config.get('vision_config', {})
704
+ has_vision = isinstance(vision, dict) and bool(vision)
705
+
706
+ return ModelMemorySpec(
707
+ repo_id=repo_id,
708
+ family='gemma4',
709
+ architecture='hybrid_gemma4_sliding_plus_full_attention',
710
+ max_position_embeddings=int(text['max_position_embeddings']),
711
+ text_hidden_size=int(text['hidden_size']),
712
+ num_hidden_layers=int(text['num_hidden_layers']),
713
+ layer_types=layer_types,
714
+ weight_footprint=weight_footprint,
715
+ cache_groups=(sliding_group, full_group),
716
+ has_vision=has_vision,
717
+ notes=(
718
+ 'Parsed as Gemma 4 hybrid sliding/full-attention model.',
719
+ f'Shared KV tail layers: {shared_tail}.',
720
+ ),
721
+ )
722
+
723
+
724
+ def load_model_spec(repo_id: str, token: str | None = None) -> ModelMemorySpec:
725
+ raw_config = _load_json(repo_id, 'config.json', token=token)
726
+ weight_footprint = _weight_footprint_from_hf(repo_id, token=token)
727
+
728
+ model_type = str(raw_config.get('model_type') or '')
729
+ text_config = raw_config.get('text_config') or raw_config
730
+ text_model_type = str(text_config.get('model_type') or model_type)
731
+
732
+ if text_model_type in {'qwen3_5_text', 'qwen3_5'}:
733
+ return _build_qwen_spec(
734
+ repo_id, raw_config, weight_footprint, family='qwen3.5'
735
+ )
736
+ if text_model_type in {'qwen3_5_moe_text', 'qwen3_5_moe'}:
737
+ return _build_qwen_spec(
738
+ repo_id, raw_config, weight_footprint, family='qwen3.6'
739
+ )
740
+ if text_model_type in {'gemma4_text', 'gemma4'}:
741
+ return _build_gemma4_spec(repo_id, raw_config, weight_footprint)
742
+
743
+ raise ValueError(
744
+ f'Unsupported model_type/text_model_type for this estimator: {model_type} / {text_model_type}'
745
+ )
746
+
747
+
748
+ def _request_floor(
749
+ model: ModelMemorySpec, deployment: DeploymentSpec
750
+ ) -> RequestFloor:
751
+ token_slope_cluster = 0.0
752
+ fixed_cluster = 0.0
753
+ notes: list[str] = []
754
+ for group in model.cache_groups:
755
+ token_slope_cluster += group.single_sequence_token_slope_cluster(
756
+ deployment
757
+ )
758
+ fixed_cluster += group.single_sequence_fixed_bytes_cluster(deployment)
759
+ notes.append(group.notes)
760
+ total_cluster = deployment.concurrent_sequences * (
761
+ token_slope_cluster * deployment.total_sequence_tokens + fixed_cluster
762
+ )
763
+ per_gpu_total = total_cluster / deployment.tensor_parallel_size
764
+ return RequestFloor(
765
+ total_bytes_per_gpu=per_gpu_total,
766
+ token_slope_bytes_per_gpu=(
767
+ token_slope_cluster / deployment.tensor_parallel_size
768
+ ),
769
+ fixed_bytes_per_gpu=(
770
+ fixed_cluster
771
+ * deployment.concurrent_sequences
772
+ / deployment.tensor_parallel_size
773
+ ),
774
+ notes=tuple(notes),
775
+ )
776
+
777
+
778
+ def _startup_fit(
779
+ model: ModelMemorySpec, deployment: DeploymentSpec
780
+ ) -> StartupFit:
781
+ weight_per_gpu: float | None = None
782
+ if model.weight_footprint.total_bytes is not None:
783
+ weight_per_gpu = (
784
+ model.weight_footprint.total_bytes / deployment.tensor_parallel_size
785
+ )
786
+ overhead = deployment.startup_overhead_policy.estimate(model, deployment)
787
+ if (
788
+ weight_per_gpu is None
789
+ or deployment.managed_budget_bytes_per_gpu is None
790
+ ):
791
+ return StartupFit(
792
+ weight_bytes_per_gpu=weight_per_gpu,
793
+ overhead_low_bytes_per_gpu=overhead.low_bytes_per_gpu,
794
+ overhead_high_bytes_per_gpu=overhead.high_bytes_per_gpu,
795
+ used_low_bytes_per_gpu=None,
796
+ used_high_bytes_per_gpu=None,
797
+ margin_low_bytes_per_gpu=None,
798
+ margin_high_bytes_per_gpu=None,
799
+ status='maybe',
800
+ )
801
+
802
+ used_low = weight_per_gpu + overhead.low_bytes_per_gpu
803
+ used_high = weight_per_gpu + overhead.high_bytes_per_gpu
804
+ margin_high = deployment.managed_budget_bytes_per_gpu - used_low
805
+ margin_low = deployment.managed_budget_bytes_per_gpu - used_high
806
+
807
+ if margin_low >= 0:
808
+ status: Literal['yes', 'no', 'maybe'] = 'yes'
809
+ elif margin_high < 0:
810
+ status = 'no'
811
+ else:
812
+ status = 'maybe'
813
+
814
+ return StartupFit(
815
+ weight_bytes_per_gpu=weight_per_gpu,
816
+ overhead_low_bytes_per_gpu=overhead.low_bytes_per_gpu,
817
+ overhead_high_bytes_per_gpu=overhead.high_bytes_per_gpu,
818
+ used_low_bytes_per_gpu=used_low,
819
+ used_high_bytes_per_gpu=used_high,
820
+ margin_low_bytes_per_gpu=margin_low,
821
+ margin_high_bytes_per_gpu=margin_high,
822
+ status=status,
823
+ )
824
+
825
+
826
+ def _steady_state_capacity(
827
+ model: ModelMemorySpec,
828
+ deployment: DeploymentSpec,
829
+ startup_fit: StartupFit,
830
+ request_floor: RequestFloor,
831
+ ) -> CapacityEstimate:
832
+ if (
833
+ deployment.managed_budget_bytes_per_gpu is None
834
+ or startup_fit.used_low_bytes_per_gpu is None
835
+ or startup_fit.used_high_bytes_per_gpu is None
836
+ ):
837
+ return CapacityEstimate(None, None, None, None, None, None, None, None)
838
+
839
+ kv_budget_low = (
840
+ deployment.managed_budget_bytes_per_gpu
841
+ - startup_fit.used_high_bytes_per_gpu
842
+ )
843
+ kv_budget_high = (
844
+ deployment.managed_budget_bytes_per_gpu
845
+ - startup_fit.used_low_bytes_per_gpu
846
+ )
847
+
848
+ if kv_budget_high <= 0:
849
+ return CapacityEstimate(
850
+ kv_budget_low, kv_budget_high, 0.0, 0.0, 0.0, 0.0, False, False
851
+ )
852
+
853
+ per_request = request_floor.total_bytes_per_gpu
854
+ if per_request <= 0:
855
+ max_conc_low = math.inf
856
+ max_conc_high = math.inf
857
+ else:
858
+ max_conc_low = max(0.0, kv_budget_low / per_request)
859
+ max_conc_high = max(0.0, kv_budget_high / per_request)
860
+
861
+ per_token_per_gpu = request_floor.token_slope_bytes_per_gpu
862
+ if per_token_per_gpu > 0:
863
+ kv_tokens_low_cluster = max(0.0, kv_budget_low / per_token_per_gpu)
864
+ kv_tokens_high_cluster = max(0.0, kv_budget_high / per_token_per_gpu)
865
+ else:
866
+ kv_tokens_low_cluster = None
867
+ kv_tokens_high_cluster = None
868
+
869
+ target = deployment.concurrent_sequences
870
+ fits_low = max_conc_low >= target
871
+ fits_high = max_conc_high >= target
872
+
873
+ return CapacityEstimate(
874
+ kv_budget_low_bytes_per_gpu=kv_budget_low,
875
+ kv_budget_high_bytes_per_gpu=kv_budget_high,
876
+ kv_tokens_low_cluster=kv_tokens_low_cluster,
877
+ kv_tokens_high_cluster=kv_tokens_high_cluster,
878
+ max_concurrency_low=max_conc_low,
879
+ max_concurrency_high=max_conc_high,
880
+ fits_target_low=fits_low,
881
+ fits_target_high=fits_high,
882
+ )
883
+
884
+
885
+ def estimate_memory(
886
+ model: ModelMemorySpec, deployment: DeploymentSpec
887
+ ) -> MemoryEstimate:
888
+ request_floor = _request_floor(model, deployment)
889
+ startup_fit = _startup_fit(model, deployment)
890
+ capacity = _steady_state_capacity(
891
+ model, deployment, startup_fit, request_floor
892
+ )
893
+ return MemoryEstimate(
894
+ model=model,
895
+ deployment=deployment,
896
+ request_floor=request_floor,
897
+ startup_fit=startup_fit,
898
+ capacity=capacity,
899
+ )
900
+
901
+
902
+ def standard_deployments(
903
+ model: ModelMemorySpec,
904
+ *,
905
+ gpu_gib: int | None = 96,
906
+ gpu_memory_utilization: float = 0.95,
907
+ language_model_only: bool = True,
908
+ ) -> list[DeploymentSpec]:
909
+ native_tokens = model.max_position_embeddings
910
+ gpu_bytes = (gpu_gib * GiB) if gpu_gib is not None else None
911
+ return [
912
+ DeploymentSpec(
913
+ name='tp4_ctx_98k',
914
+ tensor_parallel_size=4,
915
+ prompt_text_tokens=min(98304, native_tokens),
916
+ max_new_tokens=0,
917
+ gpu_memory_bytes=gpu_bytes,
918
+ gpu_memory_utilization=gpu_memory_utilization,
919
+ language_model_only=language_model_only,
920
+ ),
921
+ DeploymentSpec(
922
+ name='tp4_ctx_128k',
923
+ tensor_parallel_size=4,
924
+ prompt_text_tokens=min(131072, native_tokens),
925
+ max_new_tokens=0,
926
+ gpu_memory_bytes=gpu_bytes,
927
+ gpu_memory_utilization=gpu_memory_utilization,
928
+ language_model_only=language_model_only,
929
+ ),
930
+ DeploymentSpec(
931
+ name='tp4_ctx_256k',
932
+ tensor_parallel_size=4,
933
+ prompt_text_tokens=min(262144, native_tokens),
934
+ max_new_tokens=0,
935
+ gpu_memory_bytes=gpu_bytes,
936
+ gpu_memory_utilization=gpu_memory_utilization,
937
+ language_model_only=language_model_only,
938
+ ),
939
+ ]
940
+
941
+
942
+ def _fmt_fit_interval(low: bool | None, high: bool | None) -> str:
943
+ if low is True and high is True:
944
+ return fit_status_markup('yes')
945
+ if low is False and high is False:
946
+ return fit_status_markup('no')
947
+ return fit_status_markup('maybe')
948
+
949
+
950
+ def render_summary_matrix(
951
+ estimates: Sequence[MemoryEstimate], console: Console | None = None
952
+ ) -> None:
953
+ console = console or Console()
954
+ table = Table(
955
+ title='Deterministic startup + capacity estimate',
956
+ box=box.MINIMAL_DOUBLE_HEAD,
957
+ row_styles=['', 'dim'],
958
+ )
959
+ table.add_column('model')
960
+ table.add_column('deployment')
961
+ table.add_column('gpu GiB', justify='right')
962
+ table.add_column('managed', justify='right')
963
+ table.add_column('weights', justify='right')
964
+ table.add_column('startup ovhd', justify='right')
965
+ table.add_column('startup', justify='center')
966
+ table.add_column('req-cache', justify='right')
967
+ table.add_column('kv budget', justify='right')
968
+ table.add_column('fit@ctx', justify='center')
969
+ table.add_column('max conc', justify='right')
970
+
971
+ for est in estimates:
972
+ dep = est.deployment
973
+ sf = est.startup_fit
974
+ cap = est.capacity
975
+ row_style = None
976
+ if sf.status == 'no':
977
+ row_style = 'red'
978
+ table.add_row(
979
+ est.model.repo_id,
980
+ dep.name,
981
+ f'{dep.gpu_memory_gib:.1f}'
982
+ if dep.gpu_memory_gib is not None
983
+ else '-',
984
+ f'{gib(dep.managed_budget_bytes_per_gpu):.2f}'
985
+ if dep.managed_budget_bytes_per_gpu is not None
986
+ else '-',
987
+ gib_text(sf.weight_bytes_per_gpu),
988
+ range_text(
989
+ sf.overhead_low_bytes_per_gpu, sf.overhead_high_bytes_per_gpu
990
+ ),
991
+ fit_status_markup(sf.status),
992
+ gib_text(est.request_floor.total_bytes_per_gpu),
993
+ range_text(
994
+ cap.kv_budget_low_bytes_per_gpu or 0.0,
995
+ cap.kv_budget_high_bytes_per_gpu or 0.0,
996
+ )
997
+ if cap.kv_budget_low_bytes_per_gpu is not None
998
+ else '-',
999
+ _fmt_fit_interval(cap.fits_target_low, cap.fits_target_high),
1000
+ f'{cap.max_concurrency_low:.2f}–{cap.max_concurrency_high:.2f}'
1001
+ if cap.max_concurrency_low is not None
1002
+ else '-',
1003
+ style=row_style,
1004
+ )
1005
+ console.print(table)
1006
+
1007
+
1008
+ def render_detailed_tables(
1009
+ estimates: Sequence[MemoryEstimate], console: Console | None = None
1010
+ ) -> None:
1011
+ console = console or Console()
1012
+ for est in estimates:
1013
+ dep = est.deployment
1014
+ sf = est.startup_fit
1015
+ cap = est.capacity
1016
+ header = (
1017
+ f'family={est.model.family} | arch={est.model.architecture} | '
1018
+ f'seq={dep.total_sequence_tokens:,} | tp={dep.tensor_parallel_size} | '
1019
+ f'lm_only={dep.language_model_only} | kv_dtype={dep.kv_cache_dtype.name}'
1020
+ )
1021
+ console.print(
1022
+ Panel(
1023
+ header, title=f'{est.model.repo_id} — {dep.name}', expand=False
1024
+ )
1025
+ )
1026
+
1027
+ t = Table(box=box.SIMPLE_HEAVY)
1028
+ t.add_column('metric')
1029
+ t.add_column('value', justify='right')
1030
+ t.add_column('notes')
1031
+ t.add_row(
1032
+ 'gpu_budget/GPU',
1033
+ f'{dep.gpu_memory_gib:.2f}'
1034
+ if dep.gpu_memory_gib is not None
1035
+ else '-',
1036
+ 'Physical GPU memory',
1037
+ )
1038
+ t.add_row(
1039
+ 'managed_budget/GPU',
1040
+ gib_text(dep.managed_budget_bytes_per_gpu),
1041
+ 'gpu_memory_utilization * gpu_budget',
1042
+ )
1043
+ t.add_row(
1044
+ 'weights/GPU',
1045
+ gib_text(sf.weight_bytes_per_gpu),
1046
+ est.model.weight_footprint.notes,
1047
+ )
1048
+ t.add_row(
1049
+ 'startup_overhead/GPU',
1050
+ range_text(
1051
+ sf.overhead_low_bytes_per_gpu, sf.overhead_high_bytes_per_gpu
1052
+ ),
1053
+ dep.startup_overhead_policy.estimate(est.model, dep).notes,
1054
+ )
1055
+ t.add_row(
1056
+ 'startup_fit',
1057
+ fit_status_markup(sf.status),
1058
+ 'Compares weights + startup overhead against managed budget',
1059
+ )
1060
+ t.add_row(
1061
+ 'request_floor_cache/GPU',
1062
+ gib_text(est.request_floor.total_bytes_per_gpu),
1063
+ 'Deterministic bytes required for one configured request at this context',
1064
+ )
1065
+ t.add_row(
1066
+ 'token_slope/GPU',
1067
+ gib_text(est.request_floor.token_slope_bytes_per_gpu),
1068
+ 'Per-token request-floor cache coefficient',
1069
+ )
1070
+ t.add_row(
1071
+ 'fixed_cache/GPU',
1072
+ gib_text(est.request_floor.fixed_bytes_per_gpu),
1073
+ 'Per-sequence fixed recurrent/conv state',
1074
+ )
1075
+ if cap.kv_budget_low_bytes_per_gpu is not None and cap.kv_budget_high_bytes_per_gpu is not None:
1076
+ t.add_row(
1077
+ 'kv_budget/GPU',
1078
+ range_text(
1079
+ cap.kv_budget_low_bytes_per_gpu,
1080
+ cap.kv_budget_high_bytes_per_gpu,
1081
+ ),
1082
+ 'Managed budget minus startup-used interval',
1083
+ )
1084
+ if cap.kv_tokens_low_cluster is not None:
1085
+ t.add_row(
1086
+ 'kv_tokens/cluster',
1087
+ f'{cap.kv_tokens_low_cluster:,.0f}–{cap.kv_tokens_high_cluster:,.0f}',
1088
+ 'Cluster-wide token capacity implied by per-GPU KV budget and token slope',
1089
+ )
1090
+ if cap.max_concurrency_low is not None:
1091
+ t.add_row(
1092
+ 'max_concurrency',
1093
+ f'{cap.max_concurrency_low:.2f}–{cap.max_concurrency_high:.2f}',
1094
+ f'At ctx={dep.total_sequence_tokens:,} and concurrent_sequences target={dep.concurrent_sequences}',
1095
+ )
1096
+ console.print(t)
1097
+ console.print()
1098
+
1099
+
1100
+ def parse_deployment_arg(
1101
+ spec: str, *, default_gpu_mem_util: float, default_language_model_only: bool
1102
+ ) -> DeploymentSpec:
1103
+ # Format: name,tp,prompt,max_new[,gpu_gib][,media][,kv_dtype][,gpu_mem_util][,seqs]
1104
+ parts = spec.split(',')
1105
+ if len(parts) < 4:
1106
+ raise ValueError(
1107
+ 'Deployment spec must be name,tp,prompt,max_new[,gpu_gib][,media][,kv_dtype][,gpu_mem_util][,seqs]'
1108
+ )
1109
+ name = parts[0]
1110
+ tp = int(parts[1])
1111
+ prompt = int(parts[2])
1112
+ max_new = int(parts[3])
1113
+
1114
+ gpu_gib: int | None = None
1115
+ media = 0
1116
+ kv_dtype = DTYPE_F16
1117
+ gpu_mem_util = default_gpu_mem_util
1118
+ seqs = 1
1119
+
1120
+ if len(parts) >= 5 and parts[4]:
1121
+ gpu_gib = int(parts[4])
1122
+ if len(parts) >= 6 and parts[5]:
1123
+ media = int(parts[5])
1124
+ if len(parts) >= 7 and parts[6]:
1125
+ kv_name = parts[6].lower()
1126
+ if kv_name == 'fp8':
1127
+ kv_dtype = DTYPE_FP8
1128
+ elif kv_name == 'fp32':
1129
+ kv_dtype = DTYPE_F32
1130
+ if len(parts) >= 8 and parts[7]:
1131
+ gpu_mem_util = float(parts[7])
1132
+ if len(parts) >= 9 and parts[8]:
1133
+ seqs = int(parts[8])
1134
+
1135
+ return DeploymentSpec(
1136
+ name=name,
1137
+ tensor_parallel_size=tp,
1138
+ prompt_text_tokens=prompt,
1139
+ max_new_tokens=max_new,
1140
+ media_soft_tokens=media,
1141
+ concurrent_sequences=seqs,
1142
+ gpu_memory_bytes=(gpu_gib * GiB) if gpu_gib is not None else None,
1143
+ gpu_memory_utilization=gpu_mem_util,
1144
+ kv_cache_dtype=kv_dtype,
1145
+ language_model_only=default_language_model_only,
1146
+ )
1147
+
1148
+
1149
+ def discover_default_repo_ids(token: str | None = None) -> list[str]:
1150
+ return list(DEFAULT_PRESET_REPO_IDS)
1151
+
1152
+
1153
+ def build_arg_parser() -> argparse.ArgumentParser:
1154
+ p = argparse.ArgumentParser(
1155
+ description='Deterministic startup-fit + serving-capacity estimator for recent Qwen/Gemma models.'
1156
+ )
1157
+ p.add_argument(
1158
+ 'repo_ids',
1159
+ nargs='*',
1160
+ help='HF repo IDs. If omitted, use the if-guarded preset list in the source.',
1161
+ )
1162
+ p.add_argument(
1163
+ '--list-default-models',
1164
+ action='store_true',
1165
+ help='Print the default preset repo IDs and exit.',
1166
+ )
1167
+ p.add_argument('--token', default=None, help='HF token if needed')
1168
+ p.add_argument(
1169
+ '--standard',
1170
+ action='store_true',
1171
+ help='Run the built-in standardized deployment set',
1172
+ )
1173
+ p.add_argument(
1174
+ '--deployment',
1175
+ action='append',
1176
+ default=[],
1177
+ help='Custom deployment spec: name,tp,prompt,max_new[,gpu_gib][,media][,kv_dtype][,gpu_mem_util][,seqs]',
1178
+ )
1179
+ p.add_argument(
1180
+ '--details',
1181
+ action='store_true',
1182
+ help='Show detailed per-deployment tables',
1183
+ )
1184
+ p.add_argument(
1185
+ '--gpu-gib',
1186
+ type=int,
1187
+ default=96,
1188
+ help='GPU size for standard deployments',
1189
+ )
1190
+ p.add_argument(
1191
+ '--gpu-memory-utilization',
1192
+ type=float,
1193
+ default=0.95,
1194
+ help='Managed memory fraction for standard deployments',
1195
+ )
1196
+ p.add_argument(
1197
+ '--language-model-only',
1198
+ action='store_true',
1199
+ default=True,
1200
+ help='Assume text-only serving mode',
1201
+ )
1202
+ p.add_argument(
1203
+ '--multimodal',
1204
+ dest='language_model_only',
1205
+ action='store_false',
1206
+ help='Enable multimodal resident-memory margins',
1207
+ )
1208
+ return p
1209
+
1210
+
1211
+ def main(argv: Sequence[str] | None = None) -> int:
1212
+ args = build_arg_parser().parse_args(argv)
1213
+ console = Console()
1214
+
1215
+ repo_ids = list(args.repo_ids)
1216
+ if not repo_ids or args.list_default_models:
1217
+ repo_ids = discover_default_repo_ids(token=args.token)
1218
+ if args.list_default_models:
1219
+ for repo_id in repo_ids:
1220
+ print(repo_id)
1221
+ return 0
1222
+ console.print(
1223
+ Panel(
1224
+ '\n'.join(repo_ids) if repo_ids else '(no models enabled)',
1225
+ title='Default preset model set',
1226
+ expand=False,
1227
+ )
1228
+ )
1229
+
1230
+ explicit_deployments = [
1231
+ parse_deployment_arg(
1232
+ item,
1233
+ default_gpu_mem_util=args.gpu_memory_utilization,
1234
+ default_language_model_only=args.language_model_only,
1235
+ )
1236
+ for item in args.deployment
1237
+ ]
1238
+
1239
+ estimates: list[MemoryEstimate] = []
1240
+ for repo_id in repo_ids:
1241
+ model = load_model_spec(repo_id, token=args.token)
1242
+ model_deployments: list[DeploymentSpec] = []
1243
+ if args.standard or not explicit_deployments:
1244
+ model_deployments.extend(
1245
+ standard_deployments(
1246
+ model,
1247
+ gpu_gib=args.gpu_gib,
1248
+ gpu_memory_utilization=args.gpu_memory_utilization,
1249
+ language_model_only=args.language_model_only,
1250
+ )
1251
+ )
1252
+ model_deployments.extend(explicit_deployments)
1253
+ for dep in model_deployments:
1254
+ estimates.append(estimate_memory(model, dep))
1255
+
1256
+ render_summary_matrix(estimates, console=console)
1257
+ if args.details:
1258
+ console.print()
1259
+ render_detailed_tables(estimates, console=console)
1260
+ return 0
1261
+
1262
+
1263
+ if __name__ == '__main__':
1264
+ raise SystemExit(main())