chcc-iiitd commited on
Commit
e659e20
Β·
verified Β·
1 Parent(s): 0f9f33e

Upload 2 files

Browse files
Files changed (2) hide show
  1. app.py +1969 -0
  2. requirements.txt +4 -0
app.py ADDED
@@ -0,0 +1,1969 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ # -*- coding: utf-8 -*-
3
+ """
4
+ GPU Infrastructure Recommender for AI Models
5
+ ===========================================
6
+
7
+ A comprehensive tool for estimating VRAM requirements and recommending optimal
8
+ GPU configurations for Large Language Model (LLM) deployment and training.
9
+
10
+ Author: Rudali Huidrom
11
+ Version: 3.0.0
12
+ First Written On: 08 December 2025
13
+
14
+ Overview
15
+ --------
16
+ This module provides a Gradio-based web interface that:
17
+ 1. Calculates precise VRAM requirements for LLMs based on model specifications
18
+ 2. Estimates throughput and time-to-completion for various GPU configurations
19
+ 3. Recommends cost-effective hardware solutions for both inference and training
20
+ 4. Supports multiple quantization methods, fine-tuning strategies, and frameworks
21
+
22
+ Key Features
23
+ -----------
24
+ - Automatic model resolution from HuggingFace Hub
25
+ - Support for inference (single/batched) and training (Full FT, LoRA, QLoRA)
26
+ - Empirical throughput benchmarks for major GPU families
27
+ - Multi-GPU configuration support with communication overhead modeling
28
+ - Framework-specific optimizations (vLLM, HuggingFace, TensorRT)
29
+ - Real-time cost estimation with multiple pricing tiers
30
+
31
+ Technical Approach
32
+ -----------------
33
+ The recommender uses empirically-validated formulas derived from:
34
+ - MLPerf benchmarks and vendor specifications
35
+ - Production deployment data from real-world LLM serving
36
+ - Memory profiling of training workloads
37
+ - Community-contributed performance metrics
38
+
39
+ Accuracy: Β±15-20% variance expected due to architecture-specific optimizations,
40
+ framework versions, and runtime conditions.
41
+
42
+ Dependencies
43
+ -----------
44
+ Required:
45
+ - gradio>=3.0.0: Web interface framework
46
+ - python>=3.8: Core language support
47
+
48
+ Optional:
49
+ - transformers>=4.30.0: Automatic model config resolution
50
+ - huggingface_hub>=0.16.0: HuggingFace API access for gated models
51
+
52
+ Usage
53
+ -----
54
+ python app.py
55
+
56
+ Environment Variables:
57
+ HF_TOKEN: HuggingFace API token for accessing gated models
58
+
59
+ License
60
+ -------
61
+ Copyright (c) 2025. All rights reserved.
62
+ """
63
+
64
+ # =================================================================================================
65
+ # IMPORTS
66
+ # =================================================================================================
67
+
68
+ import math
69
+ import re
70
+ import os
71
+ from typing import Dict, List, Tuple, Optional, Any
72
+ from dataclasses import dataclass
73
+ from enum import Enum
74
+ import gradio as gr
75
+
76
+ # =================================================================================================
77
+ # CONFIGURATION AND CONSTANTS
78
+ # =================================================================================================
79
+
80
+ # Authentication token for HuggingFace API access
81
+ # Set via environment variable: export HF_TOKEN="your_token_here"
82
+ HF_TOKEN = os.getenv("HF_TOKEN", "")
83
+
84
+ # =================================================================================================
85
+ # UI STYLING
86
+ # =================================================================================================
87
+
88
+ # Custom CSS for visual differentiation of recommendation tiers
89
+ # - Budget tier: Green gradient for cost-effective options
90
+ # - Runner-up tier: Blue gradient for balanced options
91
+ # - Performance tier: Purple gradient for maximum performance
92
+ CUSTOM_CSS = """
93
+ <style>
94
+ .budget-box {
95
+ background: linear-gradient(135deg, #f0fdf4 0%, #dcfce7 100%) !important;
96
+ border: 2px solid #22c55e !important;
97
+ border-radius: 12px !important;
98
+ padding: 16px !important;
99
+ box-shadow: 0 2px 8px rgba(34, 197, 94, 0.1) !important;
100
+ }
101
+ .runner-box {
102
+ background: linear-gradient(135deg, #eff6ff 0%, #dbeafe 100%) !important;
103
+ border: 2px solid #3b82f6 !important;
104
+ border-radius: 12px !important;
105
+ padding: 16px !important;
106
+ box-shadow: 0 2px 8px rgba(59, 130, 246, 0.1) !important;
107
+ }
108
+ .perf-box {
109
+ background: linear-gradient(135deg, #faf5ff 0%, #f3e8ff 100%) !important;
110
+ border: 2px solid #a855f7 !important;
111
+ border-radius: 12px !important;
112
+ padding: 16px !important;
113
+ box-shadow: 0 2px 8px rgba(168, 85, 247, 0.1) !important;
114
+ }
115
+ .warning-box {
116
+ background-color: #fef3c7 !important;
117
+ border: 1px solid #f59e0b !important;
118
+ border-radius: 8px !important;
119
+ padding: 12px !important;
120
+ margin: 8px 0 !important;
121
+ }
122
+ .error-box {
123
+ background-color: #fee2e2 !important;
124
+ border: 1px solid #ef4444 !important;
125
+ border-radius: 8px !important;
126
+ padding: 12px !important;
127
+ margin: 8px 0 !important;
128
+ }
129
+ </style>
130
+ """
131
+
132
+ # =================================================================================================
133
+ # DATA STRUCTURES AND TYPE DEFINITIONS
134
+ # =================================================================================================
135
+
136
+ class Task(Enum):
137
+ """
138
+ Enumeration of supported computational tasks.
139
+
140
+ Attributes:
141
+ INFERENCE: Model inference/serving workloads
142
+ TRAINING: Model training/fine-tuning workloads
143
+ """
144
+ INFERENCE = "Inference"
145
+ TRAINING = "Training"
146
+
147
+ class FineTuningMethod(Enum):
148
+ """
149
+ Enumeration of supported fine-tuning strategies.
150
+
151
+ Attributes:
152
+ FULL: Full fine-tuning (all parameters trainable)
153
+ LORA: Low-Rank Adaptation (parameter-efficient, full precision base)
154
+ QLORA: Quantized LoRA (parameter-efficient, quantized base)
155
+ """
156
+ FULL = "Full Fine-Tuning"
157
+ LORA = "LoRA"
158
+ QLORA = "QLoRA"
159
+
160
+ class Framework(Enum):
161
+ """
162
+ Enumeration of supported inference/training frameworks.
163
+
164
+ Attributes:
165
+ VLLM: vLLM (optimized for high-throughput inference)
166
+ HUGGINGFACE: HuggingFace Transformers (general-purpose)
167
+ """
168
+ VLLM = "vllm"
169
+ HUGGINGFACE = "huggingface"
170
+
171
+ @dataclass
172
+ class GPUConfig:
173
+ """
174
+ Configuration specification for GPU hardware.
175
+
176
+ This dataclass encapsulates all relevant specifications and pricing
177
+ information for a GPU configuration, supporting both single and
178
+ multi-GPU setups.
179
+
180
+ Attributes:
181
+ name (str): Human-readable identifier (e.g., "Nvidia H100 SXM (8x)")
182
+ vram (int): Total VRAM across all GPUs in GB
183
+ count (int): Number of GPUs in this configuration
184
+ tflops (float): Total TFLOPS (FP16) across all GPUs
185
+ bandwidth (int): Total memory bandwidth in GB/s
186
+ price_od (float): On-demand hourly rate in INR
187
+ price_1m (float): 1-month reserved hourly rate in INR
188
+ price_6m (float): 6-month reserved hourly rate in INR
189
+ price_12m (float): 12-month reserved hourly rate in INR
190
+
191
+ Properties:
192
+ vram_per_gpu (float): VRAM per individual GPU
193
+
194
+ Methods:
195
+ get_price(tier): Returns price for specified tier
196
+ """
197
+ name: str
198
+ vram: int # Total VRAM in GB
199
+ count: int # Number of GPUs in config
200
+ tflops: float
201
+ bandwidth: int # GB/s
202
+ price_od: float # On-demand price in INR/hour
203
+ price_1m: float # 1-month reserved
204
+ price_6m: float # 6-month reserved
205
+ price_12m: float # 12-month reserved
206
+
207
+ @property
208
+ def vram_per_gpu(self) -> float:
209
+ """
210
+ Calculate VRAM per individual GPU.
211
+
212
+ Returns:
213
+ float: VRAM in GB for a single GPU in this configuration
214
+ """
215
+ return self.vram / self.count
216
+
217
+ def get_price(self, tier: str) -> float:
218
+ """
219
+ Retrieve price for specified pricing tier.
220
+
221
+ Args:
222
+ tier (str): Pricing tier ("On Demand", "1 Month Reserved", etc.)
223
+
224
+ Returns:
225
+ float: Hourly rate in INR for the specified tier
226
+ """
227
+ price_map = {
228
+ "On Demand": self.price_od,
229
+ "1 Month Reserved": self.price_1m,
230
+ "6 Month Reserved": self.price_6m,
231
+ "12 Month Reserved": self.price_12m,
232
+ }
233
+ return price_map.get(tier, self.price_od)
234
+
235
+ @dataclass
236
+ class ModelSpec:
237
+ """
238
+ Specification of transformer model architecture.
239
+
240
+ Encapsulates key architectural parameters required for accurate
241
+ memory and performance estimation.
242
+
243
+ Attributes:
244
+ params (int): Total number of model parameters
245
+ layers (int): Number of transformer layers
246
+ heads (int): Number of attention heads
247
+ kv_heads (int): Number of key/value heads (for GQA/MQA)
248
+ head_dim (int): Dimension of each attention head
249
+ context (int): Maximum context length (position embeddings)
250
+
251
+ Properties:
252
+ params_bn (float): Parameters in billions
253
+
254
+ Notes:
255
+ For standard Multi-Head Attention: kv_heads = heads
256
+ For Grouped Query Attention (GQA): kv_heads < heads
257
+ Example: Llama 3 uses heads=32, kv_heads=8 (4:1 ratio)
258
+ """
259
+ params: int # Total parameters
260
+ layers: int
261
+ heads: int
262
+ kv_heads: int
263
+ head_dim: int
264
+ context: int
265
+
266
+ @property
267
+ def params_bn(self) -> float:
268
+ """
269
+ Convert parameter count to billions.
270
+
271
+ Returns:
272
+ float: Number of parameters in billions (1e9)
273
+ """
274
+ return self.params / 1e9
275
+
276
+ # =================================================================================================
277
+ # MODEL DATABASE AND CONSTANTS
278
+ # =================================================================================================
279
+
280
+ # Popular pre-trained models available in the dropdown selector
281
+ # Sourced from HuggingFace Hub's most-used instruction-tuned models
282
+ MODEL_CHOICES = [
283
+ "meta-llama/Llama-3.3-70B-Instruct",
284
+ "meta-llama/Llama-3.1-405B-Instruct",
285
+ "meta-llama/Llama-3.1-70B-Instruct",
286
+ "meta-llama/Llama-3.1-8B-Instruct",
287
+ "meta-llama/Llama-3.2-3B-Instruct",
288
+ "meta-llama/Llama-3.2-1B-Instruct",
289
+ "Qwen/Qwen2.5-72B-Instruct",
290
+ "Qwen/Qwen2.5-32B-Instruct",
291
+ "Qwen/Qwen2.5-14B-Instruct",
292
+ "Qwen/Qwen2.5-7B-Instruct",
293
+ "Qwen/Qwen2.5-3B-Instruct",
294
+ "Qwen/Qwen2.5-1.5B-Instruct",
295
+ "Qwen/Qwen2.5-Coder-32B-Instruct",
296
+ "mistralai/Mistral-Large-Instruct-2411",
297
+ "mistralai/Mistral-Small-Instruct-2409",
298
+ "mistralai/Mistral-Nemo-Instruct-2407",
299
+ "mistralai/Mistral-7B-Instruct-v0.3",
300
+ "mistralai/Mixtral-8x22B-Instruct-v0.1",
301
+ "mistralai/Ministral-8B-Instruct-2410",
302
+ ]
303
+
304
+ # =================================================================================================
305
+ # PRECISION AND QUANTIZATION SPECIFICATIONS
306
+ # =================================================================================================
307
+
308
+ # Mapping of precision formats to bytes per parameter
309
+ # Used for accurate memory footprint calculation across different quantization schemes
310
+ #
311
+ # Precision Format Categories:
312
+ # Full Precision: fp32 (4 bytes) - Maximum accuracy, highest memory
313
+ # Half Precision: fp16, bf16 (2 bytes) - Standard training/inference
314
+ # Quantized: int8 (1 byte) - 4x compression, minimal quality loss
315
+ # Low-bit: int4, nf4 (0.5-0.56 bytes) - 8x compression, some quality degradation
316
+ # Compressed: awq, gptq (~0.52 bytes) - Advanced quantization with lookup tables
317
+ #
318
+ # Note: nf4 (NormalFloat4) is specifically designed for QLoRA and provides
319
+ # better quality than standard int4 at the same bitwidth
320
+ PRECISION_MAP = {
321
+ "float32": 4.0,
322
+ "fp32": 4.0,
323
+ "bf16": 2.0, # BFloat16 - preferred for training (better range than fp16)
324
+ "fp16": 2.0, # Float16 - standard for inference
325
+ "nf4": 0.5625, # NormalFloat4 - QLoRA's quantization format
326
+ "4bit": 0.5625,
327
+ "int4": 0.50,
328
+ "int8": 1.0,
329
+ "awq": 0.52, # Activation-aware Weight Quantization (inference-only)
330
+ "gptq": 0.52, # GPTQ quantization (inference-only)
331
+ }
332
+
333
+ # Framework-specific memory overhead (in GB)
334
+ # Represents additional memory required by the framework runtime beyond model weights
335
+ #
336
+ # Factors contributing to overhead:
337
+ # - Kernel workspace and temporary buffers
338
+ # - Execution graph and operator metadata
339
+ # - Memory pools and allocator overhead
340
+ # - Framework-specific data structures
341
+ #
342
+ # These values are empirically determined from profiling real deployments
343
+ FRAMEWORK_OVERHEAD = {
344
+ "vllm": 1.5, # PagedAttention + continuous batching optimizations
345
+ "huggingface": 3.5, # Flexible abstractions + dynamic computation graph
346
+ "tensorrt": 1.0, # Highly optimized CUDA graphs + operator fusion
347
+ }
348
+
349
+ # Quantization methods that only support inference workloads
350
+ # These methods modify weight representation in ways incompatible with gradient computation
351
+ # Training requires full-precision gradients for optimizer updates
352
+ INFERENCE_ONLY_QUANT = ['awq', 'gptq', 'exl2']
353
+
354
+ # Throughput speedup factors for different quantization methods
355
+ # Values represent throughput multiplier relative to FP16 baseline
356
+ # Based on NVIDIA TensorRT-LLM, vLLM, and MLPerf benchmarks
357
+ # Conservative estimates to avoid over-promising
358
+ QUANTIZATION_SPEEDUP = {
359
+ "fp32": 0.8, # Slightly slower than FP16 (more compute required)
360
+ "float32": 0.8,
361
+ "fp16": 1.0, # Baseline reference
362
+ "bf16": 1.0, # Same throughput as FP16
363
+ "int8": 1.8, # ~2x faster (INT8 Tensor Cores + less bandwidth)
364
+ "int4": 3.0, # ~3-4x faster (INT4 Tensor Cores + 4x less bandwidth)
365
+ "4bit": 3.0,
366
+ "nf4": 3.0, # Similar to INT4
367
+ "awq": 3.2, # Optimized INT4 quantization
368
+ "gptq": 3.2, # Optimized INT4 quantization
369
+ }
370
+
371
+ # Framework efficiency multipliers relative to vLLM baseline
372
+ # Based on production benchmarks and community reports
373
+ # vLLM is set as baseline (1.0) as it's highly optimized for inference
374
+ FRAMEWORK_SPEEDUP = {
375
+ "vllm": 1.0, # Baseline (PagedAttention, continuous batching, optimized)
376
+ "huggingface": 0.7, # More flexible but less optimized (~30% slower)
377
+ "tensorrt": 1.3, # Most optimized for NVIDIA GPUs (~30% faster)
378
+ }
379
+
380
+ # =================================================================================================
381
+ # GPU HARDWARE DATABASE
382
+ # =================================================================================================
383
+
384
+ # Comprehensive database of available GPU configurations
385
+ # Each entry represents a specific hardware configuration with associated pricing
386
+ # Pricing is in Indian Rupees (INR) per hour for various reservation tiers
387
+ GPU_DATABASE = [
388
+ # AMD MI300X
389
+ GPUConfig('AMD MI300X (1x)', 192, 1, 1300.0, 5300, 168.224, 165.048, 161.88, 148.0),
390
+ GPUConfig('AMD MI300X (2x)', 384, 2, 2600.0, 10600, 378.504, 371.358, 364.23, 333.0),
391
+ GPUConfig('AMD MI300X (4x)', 768, 4, 5200.0, 21200, 757.008, 742.716, 728.46, 666.0),
392
+ GPUConfig('AMD MI300X (8x)', 1536, 8, 10400.0, 42400, 1416.56, 1389.904, 1363.2, 1336.0),
393
+
394
+ # AMD MI325X
395
+ GPUConfig('AMD MI325X (1x)', 256, 1, 1300.0, 6000, 169.2, 123.3, 102.6, 85.5),
396
+ GPUConfig('AMD MI325X (2x)', 512, 2, 2600.0, 12000, 338.4, 246.6, 205.2, 171.0),
397
+ GPUConfig('AMD MI325X (4x)', 1024, 4, 5200.0, 24000, 676.8, 493.2, 410.4, 342.0),
398
+ GPUConfig('AMD MI325X (8x)', 2048, 8, 10400.0, 48000, 1351.8, 990.0, 820.8, 684.0),
399
+
400
+ # NVIDIA H100 SXM
401
+ GPUConfig('Nvidia H100 SXM (1x)', 80, 1, 1979.0, 3350, 153.0, 134.1, 125.1, 117.0),
402
+ GPUConfig('Nvidia H100 SXM (2x)', 160, 2, 3958.0, 6700, 306.0, 268.2, 250.2, 234.0),
403
+ GPUConfig('Nvidia H100 SXM (4x)', 320, 4, 7916.0, 13400, 612.0, 536.4, 500.4, 468.0),
404
+ GPUConfig('Nvidia H100 SXM (8x)', 640, 8, 15832.0, 26800, 1224.0, 1072.8, 1000.8, 936.0),
405
+
406
+ # NVIDIA H100 NVL
407
+ GPUConfig('Nvidia H100 NVL (1x)', 94, 1, 1671.0, 3900, 140.0, 135.0, 118.0, 100.0),
408
+ GPUConfig('Nvidia H100 NVL (2x)', 188, 2, 3342.0, 7800, 337.48, 294.44, 274.04, 257.08),
409
+ GPUConfig('Nvidia H100 NVL (4x)', 376, 4, 6684.0, 15600, 674.96, 588.88, 548.08, 514.16),
410
+ GPUConfig('Nvidia H100 NVL (8x)', 752, 8, 13368.0, 31200, 1349.92, 1177.76, 1096.16, 1028.32),
411
+
412
+ # NVIDIA H100 PCIe
413
+ GPUConfig('Nvidia H100 PCIe (1x)', 80, 1, 1513.0, 2000, 252.0, 234.0, 209.0, 185.0),
414
+ GPUConfig('Nvidia H100 PCIe (8x)', 640, 8, 12104.0, 16000, 2008.0, 1864.0, 1664.0, 1472.0),
415
+
416
+ # NVIDIA H200 SXM
417
+ GPUConfig('Nvidia H200 SXM (1x)', 141, 1, 1979.0, 4800, 140.0, 135.0, 118.0, 100.0),
418
+ GPUConfig('Nvidia H200 SXM (2x)', 282, 2, 3958.0, 9600, 510.0, 448.0, 418.0, 390.0),
419
+ GPUConfig('Nvidia H200 SXM (4x)', 564, 4, 7916.0, 19200, 1020.0, 896.0, 836.0, 780.0),
420
+ GPUConfig('Nvidia H200 SXM (8x)', 1128, 8, 15832.0, 38400, 1125.0, 1100.0, 945.0, 785.0),
421
+
422
+ # NVIDIA H200 NVL
423
+ GPUConfig('Nvidia H200 NVL (1x)', 141, 1, 1671.0, 3900, 146.38, 143.61, 140.85, 138.09),
424
+ GPUConfig('Nvidia H200 NVL (2x)', 282, 2, 3342.0, 7800, 292.75, 287.23, 281.7, 276.18),
425
+ GPUConfig('Nvidia H200 NVL (4x)', 564, 4, 6684.0, 15600, 585.5, 574.45, 563.41, 552.36),
426
+ GPUConfig('Nvidia H200 NVL (8x)', 1128, 8, 13368.0, 31200, 1171.0, 1148.91, 1126.81, 1104.72),
427
+
428
+ # NVIDIA H200 PCIe
429
+ GPUConfig('Nvidia H200 PCIe (8x)', 1128, 8, 13368.0, 31200, 3236.8, 2737.0, 2665.6, 2380.0),
430
+
431
+ # NVIDIA B200 SXM
432
+ GPUConfig('Nvidia B200 SXM (1x)', 180, 1, 4500.0, 8000, 323.0, 308.0, 293.0, 279.0),
433
+ GPUConfig('Nvidia B200 SXM (2x)', 360, 2, 9000.0, 16000, 646.0, 616.0, 586.0, 558.0),
434
+ GPUConfig('Nvidia B200 SXM (4x)', 720, 4, 18000.0, 32000, 1292.0, 1232.0, 1172.0, 1116.0),
435
+ GPUConfig('Nvidia B200 SXM (8x)', 1440, 8, 36000.0, 64000, 2584.0, 2464.0, 2344.0, 2232.0),
436
+
437
+ # NVIDIA A100 40GB
438
+ GPUConfig('Nvidia A100 40GB (1x)', 40, 1, 312.0, 1935, 136.0, 89.0, 85.0, 81.0),
439
+ GPUConfig('Nvidia A100 40GB (2x)', 80, 2, 624.0, 3870, 272.0, 178.0, 170.0, 162.0),
440
+ GPUConfig('Nvidia A100 40GB (4x)', 160, 4, 1248.0, 7740, 544.0, 356.0, 340.0, 324.0),
441
+ GPUConfig('Nvidia A100 40GB (8x)', 320, 8, 2496.0, 15480, 3175.66, 3175.66, 3175.66, 3175.66),
442
+
443
+ # NVIDIA A100 80GB
444
+ GPUConfig('Nvidia A100 80GB (1x)', 80, 1, 312.0, 1935, 135.9, 89.1, 85.5, 81.0),
445
+ GPUConfig('Nvidia A100 80GB (2x)', 160, 2, 624.0, 3870, 271.8, 178.2, 171.0, 162.0),
446
+ GPUConfig('Nvidia A100 80GB (4x)', 320, 4, 1248.0, 7740, 543.6, 356.4, 342.0, 324.0),
447
+ GPUConfig('Nvidia A100 80GB (8x)', 640, 8, 2496.0, 15480, 1087.2, 712.8, 684.0, 648.0),
448
+
449
+ # NVIDIA L40S
450
+ GPUConfig('Nvidia L40S (1x)', 48, 1, 733.0, 864, 67.5, 49.5, 49.5, 45.0),
451
+ GPUConfig('Nvidia L40S (2x)', 96, 2, 1466.0, 1728, 135.0, 99.0, 99.0, 90.0),
452
+ GPUConfig('Nvidia L40S (4x)', 192, 4, 2932.0, 3456, 306.0, 198.0, 198.0, 180.0),
453
+ GPUConfig('Nvidia L40S (8x)', 384, 8, 5864.0, 6912, 540.0, 396.0, 396.0, 360.0),
454
+
455
+ # NVIDIA L4
456
+ GPUConfig('Nvidia L4 (1x)', 24, 1, 242.0, 300, 45.07, 29.0, 26.75, 24.0),
457
+ GPUConfig('Nvidia L4 (2x)', 48, 2, 484.0, 600, 98.84, 58.0, 54.0, 48.0),
458
+ GPUConfig('Nvidia L4 (4x)', 96, 4, 968.0, 1200, 196.68, 116.0, 108.0, 96.0),
459
+ GPUConfig('Nvidia L4 (8x)', 192, 8, 1936.0, 2400, 510.37, 495.06, 459.34, 302.51),
460
+
461
+ # Intel Gaudi 2
462
+ GPUConfig('Intel Gaudi 2 (1x)', 96, 1, 180.0, 600, 57.6, 46.8, 39.6, 34.2),
463
+ GPUConfig('Intel Gaudi 2 (2x)', 192, 2, 360.0, 1200, 115.2, 93.6, 79.2, 68.4),
464
+ GPUConfig('Intel Gaudi 2 (4x)', 384, 4, 720.0, 2400, 230.4, 187.2, 158.4, 136.8),
465
+ GPUConfig('Intel Gaudi 2 (8x)', 768, 8, 1440.0, 4800, 460.8, 374.4, 316.8, 273.6),
466
+
467
+ # Intel Gaudi 3
468
+ GPUConfig('Intel Gaudi 3 (1x)', 128, 1, 459.0, 3600, 153.0, 134.1, 125.1, 117.0),
469
+ GPUConfig('Intel Gaudi 3 (2x)', 256, 2, 918.0, 7200, 306.0, 268.2, 250.2, 234.0),
470
+ GPUConfig('Intel Gaudi 3 (4x)', 512, 4, 1836.0, 14400, 612.0, 536.4, 500.4, 468.0),
471
+ GPUConfig('Intel Gaudi 3 (8x)', 1024, 8, 3672.0, 28800, 1224.0, 1072.8, 1000.8, 936.0),
472
+ ]
473
+
474
+ # =================================================================================================
475
+ # GPU Throughput Benchmarks (Empirical Data)
476
+ # =================================================================================================
477
+ # Based on real-world benchmarks from MLPerf, vendor data, and community testing
478
+ # Tokens per second per GPU for different model sizes
479
+ #
480
+ # Last Updated: 15 December 2025
481
+ # Sources:
482
+ # - MLPerf Training v3.1 (November 2023)
483
+ # - NVIDIA TensorRT-LLM benchmarks (Q4 2024)
484
+ # - vLLM project benchmarks (Q4 2024)
485
+ # - Community benchmarks from HuggingFace, Anyscale
486
+ #
487
+ # Note: These are approximate values. Actual performance varies based on:
488
+ # - Specific model architecture
489
+ # - Sequence length
490
+ # - Batch size
491
+ # - Framework optimizations
492
+ # - Hardware configuration
493
+ # Expect Β±15-20% variance in real-world usage
494
+
495
+ GPU_THROUGHPUT_BENCHMARKS = {
496
+ # Format: GPU_name -> {model_size -> (inference_tps_single, inference_tps_batched, training_tps)}
497
+ 'H100': {
498
+ 7: (120, 1400, 1800),
499
+ 13: (80, 950, 1200),
500
+ 70: (10, 90, 360),
501
+ 405: (2, 25, 90),
502
+ },
503
+ 'H200': {
504
+ 7: (130, 1500, 1950),
505
+ 13: (85, 1000, 1300),
506
+ 70: (11, 95, 390),
507
+ 405: (2, 27, 95),
508
+ },
509
+ 'B200': {
510
+ 7: (160, 1800, 2400),
511
+ 13: (105, 1200, 1600),
512
+ 70: (13, 115, 480),
513
+ 405: (3, 32, 120),
514
+ },
515
+ 'A100': {
516
+ 7: (80, 850, 900),
517
+ 13: (55, 580, 600),
518
+ 70: (6, 55, 180),
519
+ 405: (1, 5, 20),
520
+ },
521
+ 'L40S': {
522
+ 7: (45, 500, 700),
523
+ 13: (30, 330, 460),
524
+ 70: (2, 15, 80),
525
+ 405: (0.5, 2, 10),
526
+ },
527
+ 'L4': {
528
+ 7: (25, 280, 300),
529
+ 13: (10, 120, 150),
530
+ 70: (1, 8, 40),
531
+ 405: (0.3, 1, 5),
532
+ },
533
+ 'MI300X': {
534
+ 7: (80, 850, 1000),
535
+ 13: (55, 580, 660),
536
+ 70: (6, 55, 200),
537
+ 405: (1, 5, 20),
538
+ },
539
+ 'MI325X': {
540
+ 7: (88, 935, 1100),
541
+ 13: (60, 640, 720),
542
+ 70: (7, 60, 220),
543
+ 405: (1, 6, 22),
544
+ },
545
+ 'Gaudi2': {
546
+ 7: (50, 650, 800),
547
+ 13: (35, 450, 540),
548
+ 20: (25, 330, 400),
549
+ 70: (6, 80, 160),
550
+ 405: (0.8, 4, 15),
551
+ },
552
+ 'Gaudi3': {
553
+ 7: (70, 900, 1120),
554
+ 13: (50, 630, 760),
555
+ 20: (35, 460, 560),
556
+ 70: (8, 110, 225),
557
+ 405: (1, 5, 18),
558
+ },
559
+ }
560
+
561
+ def get_gpu_family(gpu_name: str) -> str:
562
+ """Extract GPU family from full GPU name."""
563
+ if 'H200' in gpu_name:
564
+ return 'H200'
565
+ elif 'H100' in gpu_name:
566
+ return 'H100'
567
+ elif 'B200' in gpu_name:
568
+ return 'B200'
569
+ elif 'A100' in gpu_name:
570
+ return 'A100'
571
+ elif 'L40S' in gpu_name:
572
+ return 'L40S'
573
+ elif 'L4' in gpu_name:
574
+ return 'L4'
575
+ elif 'MI325X' in gpu_name:
576
+ return 'MI325X'
577
+ elif 'MI300X' in gpu_name:
578
+ return 'MI300X'
579
+ elif 'Gaudi 3' in gpu_name or 'Gaudi3' in gpu_name:
580
+ return 'Gaudi3'
581
+ elif 'Gaudi 2' in gpu_name or 'Gaudi2' in gpu_name:
582
+ return 'Gaudi2'
583
+ return 'H100' # Default fallback
584
+
585
+ def interpolate_throughput(gpu_family: str, model_size_bn: float, task: str, batched: bool = False) -> float:
586
+ """
587
+ Interpolate throughput for a given GPU family and model size.
588
+
589
+ Args:
590
+ gpu_family: GPU family name (e.g., 'H100', 'A100')
591
+ model_size_bn: Model size in billions of parameters
592
+ task: 'Inference' or 'Training'
593
+ batched: Whether to use batched inference numbers
594
+
595
+ Returns:
596
+ Estimated tokens per second per GPU
597
+ """
598
+ if gpu_family not in GPU_THROUGHPUT_BENCHMARKS:
599
+ gpu_family = 'H100' # Fallback
600
+
601
+ benchmarks = GPU_THROUGHPUT_BENCHMARKS[gpu_family]
602
+
603
+ # Get the right metric index: (single_inf, batched_inf, training)
604
+ if task == "Inference":
605
+ metric_idx = 1 if batched else 0
606
+ else:
607
+ metric_idx = 2
608
+
609
+ # Create list of (size, throughput) tuples
610
+ benchmark_points = [(size, values[metric_idx]) for size, values in benchmarks.items()]
611
+ benchmark_points.sort()
612
+
613
+ # Find surrounding points for interpolation
614
+ for i in range(len(benchmark_points) - 1):
615
+ size1, tps1 = benchmark_points[i]
616
+ size2, tps2 = benchmark_points[i + 1]
617
+
618
+ if size1 <= model_size_bn <= size2:
619
+ # Log-linear interpolation (performance scales roughly inversely with size)
620
+ log_size = math.log(model_size_bn)
621
+ log_size1 = math.log(size1)
622
+ log_size2 = math.log(size2)
623
+
624
+ ratio = (log_size - log_size1) / (log_size2 - log_size1)
625
+ log_tps = math.log(tps1) + ratio * (math.log(tps2) - math.log(tps1))
626
+
627
+ return math.exp(log_tps)
628
+
629
+ # Extrapolate if outside range
630
+ if model_size_bn < benchmark_points[0][0]:
631
+ size, tps = benchmark_points[0]
632
+ return tps * (size / model_size_bn) ** 0.7
633
+ else:
634
+ size, tps = benchmark_points[-1]
635
+ return tps * (size / model_size_bn) ** 0.7
636
+
637
+ def get_lora_overhead_factor(
638
+ rank: Optional[int],
639
+ ft_method: Optional[str],
640
+ model_size_bn: float,
641
+ spec: ModelSpec
642
+ ) -> float:
643
+ """
644
+ Calculate throughput reduction factor due to LoRA adapters.
645
+
646
+ LoRA adds computational overhead through extra matrix multiplications:
647
+ - For each adapted layer: output = base_output + (B @ A @ input)
648
+ - Where A is (hidden_dim Γ— rank) and B is (rank Γ— hidden_dim)
649
+ - Higher rank = more computation = lower throughput
650
+
651
+ This function uses empirically-measured overhead from:
652
+ - QLoRA paper (Dettmers et al., 2023)
653
+ - Community benchmarks (HuggingFace, Axolotl)
654
+ - Production LoRA training deployments
655
+
656
+ Args:
657
+ rank: LoRA rank (None if not using LoRA/QLoRA)
658
+ ft_method: Fine-tuning method
659
+ model_size_bn: Model size in billions of parameters
660
+ spec: Model specification (for hidden_dim and layers)
661
+
662
+ Returns:
663
+ Throughput multiplier (< 1.0 means slower due to LoRA overhead)
664
+
665
+ Example:
666
+ >>> get_lora_overhead_factor(64, "QLoRA", 7.0, spec)
667
+ 0.87 # 13% throughput reduction with rank=64 on 7B model
668
+ """
669
+ if ft_method not in ["LoRA", "QLoRA"] or rank is None:
670
+ return 1.0 # No LoRA overhead
671
+
672
+ # Empirical overhead data from real-world benchmarks
673
+ # These account for: kernel launch overhead, memory bandwidth, cache effects
674
+ # Not just theoretical FLOPs (which underestimate actual impact)
675
+
676
+ if model_size_bn <= 10:
677
+ # 7B models - LoRA overhead is most significant
678
+ # Kernel launch costs and memory bandwidth dominate
679
+ overhead_map = {
680
+ 8: 0.98, # 2% slowdown
681
+ 16: 0.96, # 4% slowdown
682
+ 32: 0.92, # 8% slowdown
683
+ 64: 0.87, # 13% slowdown
684
+ 128: 0.80, # 20% slowdown
685
+ 256: 0.70, # 30% slowdown
686
+ }
687
+ elif model_size_bn <= 20:
688
+ # 13B models
689
+ overhead_map = {
690
+ 8: 0.99,
691
+ 16: 0.97,
692
+ 32: 0.94,
693
+ 64: 0.89,
694
+ 128: 0.83,
695
+ 256: 0.74,
696
+ }
697
+ elif model_size_bn <= 100:
698
+ # 70B models - LoRA overhead is less significant
699
+ # Base model computation dominates
700
+ overhead_map = {
701
+ 8: 0.99,
702
+ 16: 0.98,
703
+ 32: 0.96,
704
+ 64: 0.92,
705
+ 128: 0.87,
706
+ 256: 0.79,
707
+ }
708
+ else:
709
+ # 405B+ models - LoRA is proportionally tiny
710
+ overhead_map = {
711
+ 8: 1.00, # Negligible
712
+ 16: 0.99,
713
+ 32: 0.98,
714
+ 64: 0.95,
715
+ 128: 0.91,
716
+ 256: 0.84,
717
+ }
718
+
719
+ # Find or interpolate for the given rank
720
+ if rank in overhead_map:
721
+ return overhead_map[rank]
722
+
723
+ # Interpolate for ranks not in the map
724
+ ranks = sorted(overhead_map.keys())
725
+ for i in range(len(ranks) - 1):
726
+ if ranks[i] < rank < ranks[i+1]:
727
+ r1, r2 = ranks[i], ranks[i+1]
728
+ v1, v2 = overhead_map[r1], overhead_map[r2]
729
+ # Linear interpolation in log-space for smoother scaling
730
+ import math
731
+ log_rank = math.log(rank)
732
+ log_r1 = math.log(r1)
733
+ log_r2 = math.log(r2)
734
+ ratio = (log_rank - log_r1) / (log_r2 - log_r1)
735
+ return v1 + ratio * (v2 - v1)
736
+
737
+ # Extrapolate if beyond range
738
+ if rank < ranks[0]:
739
+ return 1.0 # No overhead for very small ranks (< 8)
740
+ else:
741
+ # For very high ranks (> 256), assume overhead continues growing
742
+ return max(0.5, overhead_map[ranks[-1]] * 0.9)
743
+
744
+ def calculate_throughput(
745
+ gpu_config: GPUConfig,
746
+ spec: ModelSpec,
747
+ task: str,
748
+ batch_size: int,
749
+ precision: str = "fp16",
750
+ framework: str = "vllm",
751
+ ft_method: Optional[str] = None,
752
+ rank: Optional[int] = None
753
+ ) -> Tuple[float, float, str]:
754
+ """
755
+ Calculate estimated throughput for a GPU configuration.
756
+
757
+ Throughput varies significantly by:
758
+ 1. Quantization (INT8/INT4 Tensor Cores provide 2-4x speedup)
759
+ 2. Framework (TensorRT-LLM > vLLM > HuggingFace)
760
+ 3. Batch size and GPU architecture
761
+ 4. LoRA rank (for training - higher rank = more overhead)
762
+
763
+ Args:
764
+ gpu_config: GPU configuration
765
+ spec: Model specification
766
+ task: 'Inference' or 'Training'
767
+ batch_size: Batch size
768
+ precision: Quantization/precision format (e.g., 'fp16', 'int8', 'int4')
769
+ framework: Inference framework ('vllm', 'huggingface', 'tensorrt')
770
+ ft_method: Fine-tuning method ('LoRA', 'QLoRA', 'Full Fine-Tuning')
771
+ rank: LoRA rank (only used for LoRA/QLoRA training)
772
+
773
+ Returns:
774
+ Tuple of (tokens_per_second_per_gpu, total_tokens_per_second, description)
775
+
776
+ Example:
777
+ >>> # INT4 with TensorRT-LLM is ~4x faster than FP16 HuggingFace
778
+ >>> calc_throughput(h100, llama7b, "Inference", 32, "int4", "tensorrt")
779
+ (3900, 3900, "Batched inference (int4, tensorrt) - 4.2x speedup")
780
+
781
+ >>> # LoRA rank affects training throughput
782
+ >>> calc_throughput(h100, llama7b, "Training", 16, "nf4", "huggingface", "QLoRA", 64)
783
+ (1600, 1600, "Training throughput (nf4, LoRA r=64, 89% efficiency)")
784
+ """
785
+ gpu_family = get_gpu_family(gpu_config.name)
786
+ model_size_bn = spec.params_bn
787
+
788
+ # Determine if we should use batched numbers
789
+ use_batched = batch_size >= 8 if task == "Inference" else False
790
+
791
+ # Get base throughput per GPU (assumes FP16 on vLLM baseline)
792
+ tps_per_gpu = interpolate_throughput(gpu_family, model_size_bn, task, use_batched)
793
+
794
+ # Apply quantization speedup multiplier
795
+ # INT8/INT4 are significantly faster due to specialized Tensor Cores
796
+ quant_speedup = QUANTIZATION_SPEEDUP.get(precision, 1.0)
797
+ tps_per_gpu *= quant_speedup
798
+
799
+ # Apply framework efficiency multiplier
800
+ # TensorRT-LLM is more optimized than vLLM, HuggingFace is less optimized
801
+ framework_speedup = FRAMEWORK_SPEEDUP.get(framework.lower(), 1.0)
802
+ tps_per_gpu *= framework_speedup
803
+
804
+ # Apply LoRA overhead if applicable (TRAINING ONLY)
805
+ # LoRA adds extra matrix multiplications that reduce throughput
806
+ lora_overhead = 1.0 # Default: no overhead
807
+ if task == "Training":
808
+ lora_overhead = get_lora_overhead_factor(rank, ft_method, model_size_bn, spec)
809
+ tps_per_gpu *= lora_overhead
810
+
811
+ # Calculate combined speedup for description
812
+ combined_speedup = quant_speedup * framework_speedup
813
+
814
+ # Apply batch scaling for inference
815
+ if task == "Inference" and batch_size > 1:
816
+ if use_batched:
817
+ # Already using batched numbers, apply efficiency factor
818
+ # Remove cap to allow larger batches to increase throughput
819
+ batch_efficiency = (batch_size / 32) ** 0.7
820
+ tps_per_gpu *= batch_efficiency
821
+ else:
822
+ # Single stream numbers, scale by batch with diminishing returns
823
+ batch_efficiency = min(1.0, (batch_size / 8) ** 0.6)
824
+ tps_per_gpu *= batch_efficiency * batch_size
825
+
826
+ # Apply batch scaling for training
827
+ if task == "Training" and batch_size > 1:
828
+ # Remove cap to allow larger batches to increase throughput
829
+ batch_efficiency = (batch_size / 8) ** 0.7
830
+ tps_per_gpu *= batch_efficiency
831
+
832
+ # Apply multi-GPU communication overhead
833
+ if gpu_config.count > 1:
834
+ if gpu_config.count <= 4:
835
+ comm_efficiency = 0.90
836
+ elif gpu_config.count <= 8:
837
+ comm_efficiency = 0.85
838
+ else:
839
+ comm_efficiency = 0.75
840
+ tps_per_gpu *= comm_efficiency
841
+
842
+ # Total throughput across all GPUs
843
+ total_tps = tps_per_gpu * gpu_config.count
844
+
845
+ # Generate description
846
+ if task == "Inference":
847
+ desc = f"{'Batched' if use_batched else 'Single-stream'} inference ({precision}, {framework})"
848
+ if combined_speedup != 1.0:
849
+ desc += f" - {combined_speedup:.1f}x speedup"
850
+ else:
851
+ desc = f"Training throughput ({precision})"
852
+ if ft_method in ["LoRA", "QLoRA"] and rank:
853
+ # Add LoRA rank info and efficiency
854
+ desc += f", LoRA r={rank}, {lora_overhead:.0%} efficiency"
855
+
856
+ if gpu_config.count > 1:
857
+ desc += f" ({gpu_config.count}x GPUs, {comm_efficiency:.0%} efficiency)"
858
+
859
+ return tps_per_gpu, total_tps, desc
860
+
861
+ def format_time_estimate(total_tokens: int, throughput_tps: float) -> str:
862
+ """
863
+ Format time estimate based on tokens and throughput.
864
+
865
+ Args:
866
+ total_tokens: Total tokens to process
867
+ throughput_tps: Throughput in tokens per second
868
+
869
+ Returns:
870
+ Formatted time string
871
+ """
872
+ if throughput_tps <= 0:
873
+ return "N/A"
874
+
875
+ seconds = total_tokens / throughput_tps
876
+
877
+ if seconds < 60:
878
+ return f"{seconds:.1f}s"
879
+ elif seconds < 3600:
880
+ return f"{seconds/60:.1f}m"
881
+ elif seconds < 86400:
882
+ return f"{seconds/3600:.1f}h"
883
+ else:
884
+ return f"{seconds/86400:.1f}d"
885
+
886
+ # =================================================================================================
887
+ # Validation and Input Processing
888
+ # =================================================================================================
889
+
890
+ class ValidationError(Exception):
891
+ """Custom exception for validation errors."""
892
+ pass
893
+
894
+ def validate_inputs(
895
+ seq_len: float,
896
+ batch: float,
897
+ rank: float,
898
+ sample_count: float,
899
+ input_tokens: float,
900
+ output_tokens: float
901
+ ) -> List[str]:
902
+ """
903
+ Validate user inputs and return list of warnings/errors.
904
+
905
+ Returns:
906
+ List of validation messages (empty if all valid)
907
+ """
908
+ warnings = []
909
+
910
+ # Check reasonable ranges
911
+ if seq_len < 128:
912
+ warnings.append("WARNING: Context length < 128 may be too small for most models")
913
+ if seq_len > 100000:
914
+ warnings.append("WARNING: Very large context length will require significant VRAM")
915
+
916
+ if batch < 1:
917
+ warnings.append("ERROR: Batch size must be at least 1")
918
+ if batch > 512:
919
+ warnings.append("WARNING: Very large batch size may exceed VRAM limits")
920
+
921
+ if rank < 4:
922
+ warnings.append("WARNING: LoRA rank < 4 may be too low for effective fine-tuning")
923
+ if rank > 256:
924
+ warnings.append("WARNING: LoRA rank > 256 may be inefficient (diminishing returns)")
925
+
926
+ if sample_count < 1:
927
+ warnings.append("ERROR: Sample count must be at least 1")
928
+
929
+ if input_tokens < 1 or output_tokens < 1:
930
+ warnings.append("ERROR: Token counts must be positive")
931
+
932
+ return warnings
933
+
934
+ # =================================================================================================
935
+ # Model Resolution
936
+ # =================================================================================================
937
+
938
+ # Check if transformers library is available
939
+ try:
940
+ from transformers import AutoConfig
941
+ from huggingface_hub import HfApi
942
+ TRANSFORMERS_AVAILABLE = True
943
+ except ImportError:
944
+ TRANSFORMERS_AVAILABLE = False
945
+ print("Tip: Install 'huggingface_hub' & 'transformers' for automatic model resolution")
946
+
947
+ def estimate_architecture(params_bn: float) -> ModelSpec:
948
+ """
949
+ Estimates model architecture based on parameter count.
950
+
951
+ Args:
952
+ params_bn: Number of parameters in billions
953
+
954
+ Returns:
955
+ ModelSpec with estimated architecture
956
+ """
957
+ # Rule-of-thumb estimates based on common architectures
958
+ if params_bn < 1:
959
+ layers, heads, kv_heads = 12, 12, 12
960
+ elif params_bn < 4:
961
+ layers, heads, kv_heads = 20, 16, 16
962
+ elif params_bn < 10:
963
+ layers, heads, kv_heads = 32, 32, 8
964
+ elif params_bn < 20:
965
+ layers, heads, kv_heads = 40, 40, 8
966
+ elif params_bn < 50:
967
+ layers, heads, kv_heads = 60, 64, 8
968
+ else:
969
+ layers, heads, kv_heads = 80, 80, 8
970
+
971
+ return ModelSpec(
972
+ params=int(params_bn * 1e9),
973
+ layers=layers,
974
+ heads=heads,
975
+ kv_heads=kv_heads,
976
+ head_dim=128,
977
+ context=32768
978
+ )
979
+
980
+ def fetch_hf_config(repo_id: str, token: Optional[str] = None) -> Tuple[ModelSpec, str]:
981
+ """
982
+ Fetches model configuration from Hugging Face Hub.
983
+
984
+ Args:
985
+ repo_id: Repository ID (e.g., 'meta-llama/Llama-2-7b')
986
+ token: Optional HuggingFace authentication token
987
+
988
+ Returns:
989
+ Tuple of (ModelSpec, repo_id)
990
+
991
+ Raises:
992
+ PermissionError: If model is gated
993
+ FileNotFoundError: If model doesn't exist
994
+ Exception: Other errors
995
+ """
996
+ if not TRANSFORMERS_AVAILABLE:
997
+ raise ImportError("transformers library not available")
998
+
999
+ try:
1000
+ config = AutoConfig.from_pretrained(repo_id, trust_remote_code=True, token=token)
1001
+
1002
+ # Get parameter count
1003
+ params = getattr(config, "num_parameters", None)
1004
+ if callable(params):
1005
+ params = params()
1006
+
1007
+ # Estimate if not available
1008
+ if params is None:
1009
+ hidden = config.hidden_size
1010
+ layers = config.num_hidden_layers
1011
+ intermediate = getattr(config, "intermediate_size", hidden * 4)
1012
+ params = layers * (4 * hidden * hidden + 3 * hidden * intermediate)
1013
+
1014
+ spec = ModelSpec(
1015
+ params=params,
1016
+ layers=config.num_hidden_layers,
1017
+ heads=config.num_attention_heads,
1018
+ kv_heads=getattr(config, "num_key_value_heads", config.num_attention_heads),
1019
+ head_dim=config.hidden_size // config.num_attention_heads,
1020
+ context=getattr(config, "max_position_embeddings", 8192)
1021
+ )
1022
+
1023
+ return spec, repo_id
1024
+
1025
+ except Exception as e:
1026
+ err = str(e).lower()
1027
+ if "401" in err or "403" in err or "gated" in err:
1028
+ raise PermissionError(f"Model '{repo_id}' is gated. Please provide HF token.")
1029
+ if "404" in err or "not found" in err:
1030
+ raise FileNotFoundError(f"Model '{repo_id}' not found on Hugging Face Hub")
1031
+ raise e
1032
+
1033
+ def resolve_model(model_name: str, token: Optional[str] = None) -> Tuple[ModelSpec, str, List[str]]:
1034
+ """
1035
+ Resolves model specification from name or parameter count.
1036
+
1037
+ Args:
1038
+ model_name: Either HF repo ID or parameter count (e.g., "7B")
1039
+ token: Optional HuggingFace token
1040
+
1041
+ Returns:
1042
+ Tuple of (ModelSpec, source_description, logs)
1043
+ """
1044
+ logs = []
1045
+
1046
+ # Try to parse as parameter count (e.g., "7B", "70B")
1047
+ match = re.match(r'^(\d+\.?\d*)\s*[Bb]', model_name.strip())
1048
+ if match:
1049
+ params_bn = float(match.group(1))
1050
+ spec = estimate_architecture(params_bn)
1051
+ logs.append(f"Using estimated architecture for {params_bn}B parameters")
1052
+ return spec, f"Estimated {params_bn}B model", logs
1053
+
1054
+ # Try to fetch from Hugging Face
1055
+ if TRANSFORMERS_AVAILABLE:
1056
+ try:
1057
+ spec, repo = fetch_hf_config(model_name, token)
1058
+ logs.append(f"Successfully loaded config from Hugging Face: {repo}")
1059
+ logs.append(f" Parameters: {spec.params_bn:.2f}B, Layers: {spec.layers}, Context: {spec.context}")
1060
+ return spec, f"HuggingFace: {repo}", logs
1061
+ except PermissionError as e:
1062
+ logs.append(f"PERMISSION DENIED: {str(e)}")
1063
+ logs.append(" Falling back to estimation...")
1064
+ except FileNotFoundError as e:
1065
+ logs.append(f"NOT FOUND: {str(e)}")
1066
+ logs.append(" Falling back to estimation...")
1067
+ except Exception as e:
1068
+ logs.append(f"ERROR: Error loading config: {str(e)}")
1069
+ logs.append(" Falling back to estimation...")
1070
+
1071
+ # Fallback: estimate from common model names
1072
+ name_lower = model_name.lower()
1073
+ if "405b" in name_lower:
1074
+ params_bn = 405
1075
+ elif "70b" in name_lower or "72b" in name_lower:
1076
+ params_bn = 70
1077
+ elif "34b" in name_lower or "32b" in name_lower:
1078
+ params_bn = 34
1079
+ elif "13b" in name_lower or "14b" in name_lower:
1080
+ params_bn = 13
1081
+ elif "7b" in name_lower or "8b" in name_lower:
1082
+ params_bn = 7
1083
+ elif "3b" in name_lower:
1084
+ params_bn = 3
1085
+ elif "1b" in name_lower or "1.5b" in name_lower:
1086
+ params_bn = 1.5
1087
+ else:
1088
+ params_bn = 7 # Default fallback
1089
+ logs.append("WARNING: Could not determine model size, using 7B as default")
1090
+
1091
+ spec = estimate_architecture(params_bn)
1092
+ logs.append(f"Using estimated architecture for {params_bn}B parameters")
1093
+
1094
+ return spec, f"Estimated {params_bn}B model", logs
1095
+
1096
+ # =================================================================================================
1097
+ # VRAM CALCULATION ENGINE
1098
+ # =================================================================================================
1099
+ #
1100
+ # This section contains the core memory estimation algorithms for transformer models.
1101
+ # All calculations are based on empirically-validated formulas derived from:
1102
+ # - Production deployments of LLMs
1103
+ # - Memory profiling of training workloads
1104
+ # - Vendor specifications and benchmarks
1105
+ # - Academic research on transformer efficiency
1106
+ #
1107
+ # Accuracy: Β±10-15% for model weights, Β±15-20% for dynamic allocations
1108
+
1109
+ def calculate_model_weights(spec: ModelSpec, precision: str) -> float:
1110
+ """
1111
+ Calculate memory footprint of model parameters.
1112
+
1113
+ Computes the storage requirement for all model parameters based on the
1114
+ specified precision format. This is the base memory requirement before
1115
+ considering any runtime allocations.
1116
+
1117
+ Formula:
1118
+ memory_gb = (num_parameters Γ— bytes_per_parameter) / (1024Β³)
1119
+
1120
+ Args:
1121
+ spec (ModelSpec): Model architecture specification containing parameter count
1122
+ precision (str): Precision format (e.g., 'fp16', 'nf4', 'int8')
1123
+ Must be a valid key in PRECISION_MAP
1124
+
1125
+ Returns:
1126
+ float: Memory requirement in gigabytes (GB)
1127
+
1128
+ Example:
1129
+ >>> spec = ModelSpec(params=7_000_000_000, ...) # 7B parameters
1130
+ >>> calculate_model_weights(spec, 'fp16')
1131
+ 13.0 # 7B Γ— 2 bytes / 1024Β³ β‰ˆ 13 GB
1132
+ """
1133
+ bytes_per_param = PRECISION_MAP.get(precision, 2.0) # Default to fp16 if unknown
1134
+ return (spec.params * bytes_per_param) / (1024**3)
1135
+
1136
+ def calculate_kv_cache(
1137
+ spec: ModelSpec,
1138
+ batch_size: int,
1139
+ seq_len: int,
1140
+ precision: str
1141
+ ) -> float:
1142
+ """
1143
+ Calculate Key-Value cache memory requirement for transformer inference.
1144
+
1145
+ The KV cache stores computed key and value vectors from attention layers
1146
+ to avoid recomputation during autoregressive generation. This is the primary
1147
+ dynamic memory component during inference.
1148
+
1149
+ Formula:
1150
+ kv_memory = 2 Γ— L Γ— B Γ— S Γ— H_kv Γ— D Γ— P
1151
+
1152
+ Where:
1153
+ 2 = Keys + Values (separate tensors)
1154
+ L = Number of layers
1155
+ B = Batch size
1156
+ S = Sequence length
1157
+ H_kv = Number of key/value heads (for GQA/MQA architectures)
1158
+ D = Dimension per head
1159
+ P = Bytes per element (precision)
1160
+
1161
+ Important Notes:
1162
+ - For standard Multi-Head Attention: H_kv = H (total heads)
1163
+ - For Grouped Query Attention (GQA): H_kv < H
1164
+ Example: Llama 3 uses H=32, H_kv=8 (4:1 ratio for efficiency)
1165
+ - KV cache typically kept at higher precision (fp16) even when model
1166
+ is quantized, as aggressive quantization degrades generation quality
1167
+
1168
+ Args:
1169
+ spec (ModelSpec): Model architecture specification
1170
+ batch_size (int): Number of sequences processed in parallel
1171
+ seq_len (int): Maximum sequence length to cache
1172
+ precision (str): Precision format for KV cache storage
1173
+
1174
+ Returns:
1175
+ float: Memory requirement in gigabytes (GB)
1176
+
1177
+ Example:
1178
+ >>> spec = ModelSpec(layers=32, kv_heads=8, head_dim=128, ...)
1179
+ >>> calculate_kv_cache(spec, batch_size=32, seq_len=2048, precision='fp16')
1180
+ 4.0 # Approximately 4 GB for this configuration
1181
+ """
1182
+ bytes_per_elem = PRECISION_MAP.get(precision, 2.0)
1183
+
1184
+ # KV cache is typically not quantized as aggressively as model weights
1185
+ # to maintain generation quality. Force fp16 for low-bit formats.
1186
+ if precision in ['nf4', '4bit', 'int4']:
1187
+ bytes_per_elem = 2.0 # Override to fp16 for quality preservation
1188
+
1189
+ kv_memory_bytes = (
1190
+ 2 # Separate K and V tensors
1191
+ * spec.layers # One cache per transformer layer
1192
+ * batch_size # Parallel sequences
1193
+ * seq_len # Tokens per sequence
1194
+ * spec.kv_heads # Key/value heads (may differ from query heads in GQA)
1195
+ * spec.head_dim # Dimension of each attention head
1196
+ * bytes_per_elem # Precision-dependent storage size
1197
+ )
1198
+
1199
+ return kv_memory_bytes / (1024**3) # Convert bytes to GB
1200
+
1201
+ def calculate_activations(
1202
+ spec: ModelSpec,
1203
+ batch_size: int,
1204
+ seq_len: int,
1205
+ precision: str,
1206
+ use_checkpointing: bool = True # Most frameworks use some form of checkpointing
1207
+ ) -> float:
1208
+ """
1209
+ Calculate activation memory for training (more accurate estimate).
1210
+
1211
+ Args:
1212
+ spec: Model specification
1213
+ batch_size: Batch size
1214
+ seq_len: Sequence length
1215
+ precision: Precision format
1216
+ use_checkpointing: Whether gradient checkpointing is used
1217
+
1218
+ Returns:
1219
+ Memory in GB
1220
+ """
1221
+ bytes_per_elem = PRECISION_MAP.get(precision, 2.0)
1222
+ hidden_size = spec.heads * spec.head_dim
1223
+
1224
+ # Activation memory depends on gradient checkpointing strategy:
1225
+ # - No checkpointing: Store all intermediate activations (~34x)
1226
+ # - Selective checkpointing: Recompute some activations (~12x)
1227
+ # Most modern frameworks use some form of checkpointing by default
1228
+ multiplier = 12 if use_checkpointing else 34
1229
+
1230
+ activation_bytes = (
1231
+ batch_size
1232
+ * seq_len
1233
+ * hidden_size
1234
+ * spec.layers
1235
+ * multiplier
1236
+ * bytes_per_elem
1237
+ )
1238
+
1239
+ return activation_bytes / (1024**3)
1240
+
1241
+ def calculate_optimizer_states(
1242
+ model_weights_gb: float,
1243
+ ft_method: str,
1244
+ rank: int,
1245
+ spec: ModelSpec
1246
+ ) -> Tuple[float, str]:
1247
+ """
1248
+ Calculate optimizer state memory.
1249
+
1250
+ Args:
1251
+ model_weights_gb: Model weights in GB
1252
+ ft_method: Fine-tuning method
1253
+ rank: LoRA rank (used for LoRA/QLoRA)
1254
+ spec: Model specification (for calculating adapter size)
1255
+
1256
+ Returns:
1257
+ Tuple of (memory in GB, description)
1258
+ """
1259
+ if ft_method == "Full Fine-Tuning":
1260
+ # Adam: momentum + variance, both stored at fp32
1261
+ # Model weights are fp16 (2 bytes), optimizer states are fp32 (4 bytes each)
1262
+ # Total: 2 states Γ— 4 bytes = 8 bytes per param vs 2 bytes for model
1263
+ # = 4x model weight size
1264
+ optimizer_gb = model_weights_gb * 4
1265
+ return optimizer_gb, "Optimizer (Adam - Full FT)"
1266
+ else:
1267
+ # LoRA/QLoRA: only optimizer states for adapter weights
1268
+ # Adapter parameters per layer: 2 matrices (A and B) of size (rank Γ— hidden_dim)
1269
+ hidden_dim = spec.heads * spec.head_dim
1270
+ adapter_params = 2 * rank * hidden_dim * spec.layers
1271
+
1272
+ # Adapter params stored at fp16
1273
+ adapter_params_gb = (adapter_params * 2) / (1024**3)
1274
+
1275
+ # Adam optimizer states at fp32: 2x params at fp32 = 4x params at fp16
1276
+ optimizer_gb = adapter_params_gb * 4
1277
+
1278
+ return optimizer_gb, f"Optimizer (LoRA r={rank})"
1279
+
1280
+ def calculate_gradients(
1281
+ model_weights_gb: float,
1282
+ ft_method: str,
1283
+ rank: int,
1284
+ spec: ModelSpec
1285
+ ) -> Tuple[float, str]:
1286
+ """
1287
+ Calculate gradient memory.
1288
+
1289
+ Args:
1290
+ model_weights_gb: Model weights in GB
1291
+ ft_method: Fine-tuning method
1292
+ rank: LoRA rank
1293
+ spec: Model specification (for calculating adapter size)
1294
+
1295
+ Returns:
1296
+ Tuple of (memory in GB, description)
1297
+ """
1298
+ if ft_method == "Full Fine-Tuning":
1299
+ # Full fine-tuning: gradients for all parameters
1300
+ # Gradients typically stored at fp32 for numerical stability
1301
+ gradient_gb = model_weights_gb * 2 # fp32 vs fp16
1302
+ return gradient_gb, "Gradients (Full FT)"
1303
+ else:
1304
+ # LoRA/QLoRA: only gradients for adapter weights
1305
+ # Calculate actual adapter size based on model architecture
1306
+ hidden_dim = spec.heads * spec.head_dim
1307
+ adapter_params = 2 * rank * hidden_dim * spec.layers
1308
+
1309
+ # Gradients stored at fp16 (same precision as adapter params)
1310
+ adapter_gradient_gb = (adapter_params * 2) / (1024**3)
1311
+
1312
+ return adapter_gradient_gb, f"Gradients (LoRA r={rank})"
1313
+
1314
+ def calculate_vram(
1315
+ spec: ModelSpec,
1316
+ precision: str,
1317
+ batch_size: int,
1318
+ seq_len: int,
1319
+ task: str,
1320
+ framework: str,
1321
+ ft_method: Optional[str] = None,
1322
+ rank: Optional[int] = None
1323
+ ) -> Tuple[float, float, float, str]:
1324
+ """
1325
+ Main VRAM calculation function.
1326
+
1327
+ Args:
1328
+ spec: Model specification
1329
+ precision: Precision format
1330
+ batch_size: Batch size
1331
+ seq_len: Sequence length
1332
+ task: 'Inference' or 'Training'
1333
+ framework: Framework being used
1334
+ ft_method: Fine-tuning method (for training)
1335
+ rank: LoRA rank (for LoRA/QLoRA)
1336
+
1337
+ Returns:
1338
+ Tuple of (total_vram_gb, weights_gb, variable_gb, variable_label, actual_precision)
1339
+ """
1340
+ # KEY DIFFERENCE: QLoRA vs LoRA vs Full FT
1341
+ # QLoRA: Base model stays quantized (e.g., 4-bit), adapters in fp16/bf16
1342
+ # LoRA: Base model in fp16/bf16, adapters in fp16/bf16
1343
+ # Full FT: Base model MUST be fp16/bf16 (all params trainable)
1344
+
1345
+ # Track the actual precision being used (may differ from user selection)
1346
+ actual_precision = precision
1347
+
1348
+ if task == "Training" and ft_method in ["LoRA", "Full Fine-Tuning"]:
1349
+ # LoRA and Full FT require full precision base model
1350
+ # Even if user selected quantization, these methods need bf16/fp16
1351
+ if precision in ['nf4', '4bit', 'int4', 'int8', 'awq', 'gptq']:
1352
+ # Override to bf16
1353
+ base_precision = 'bf16'
1354
+ actual_precision = 'bf16' # Track the override
1355
+ weights_gb = calculate_model_weights(spec, base_precision)
1356
+ else:
1357
+ weights_gb = calculate_model_weights(spec, precision)
1358
+ else:
1359
+ # QLoRA or Inference: use selected precision
1360
+ # QLoRA can use quantized base because it's frozen
1361
+ weights_gb = calculate_model_weights(spec, precision)
1362
+
1363
+ # Calculate KV cache
1364
+ kv_cache_gb = calculate_kv_cache(spec, batch_size, seq_len, precision)
1365
+
1366
+ if task == "Inference":
1367
+ # Inference: weights + KV cache + framework overhead
1368
+ overhead_gb = FRAMEWORK_OVERHEAD.get(framework, 2.0)
1369
+ variable_gb = kv_cache_gb + overhead_gb
1370
+ variable_label = f"KV Cache + {framework.upper()} Overhead"
1371
+
1372
+ else: # Training
1373
+ # Training: weights + KV + activations + optimizer + gradients
1374
+ activations_gb = calculate_activations(spec, batch_size, seq_len, precision)
1375
+ optimizer_gb, _ = calculate_optimizer_states(weights_gb, ft_method or "Full Fine-Tuning", rank or 64, spec)
1376
+ gradients_gb, _ = calculate_gradients(weights_gb, ft_method or "Full Fine-Tuning", rank or 64, spec)
1377
+
1378
+ # Add LoRA adapter weights (small additional memory)
1379
+ if ft_method in ["LoRA", "QLoRA"]:
1380
+ # LoRA adapters: A and B matrices per layer
1381
+ # Each matrix is (rank Γ— hidden_dim), stored at fp16
1382
+ hidden_dim = spec.heads * spec.head_dim
1383
+ adapter_params = 2 * rank * hidden_dim * spec.layers
1384
+ adapter_gb = (adapter_params * 2) / (1024**3) # fp16
1385
+
1386
+ variable_gb = kv_cache_gb + activations_gb + optimizer_gb + gradients_gb + adapter_gb
1387
+ variable_label = "KV + Activations + Optimizer + Gradients + LoRA Adapters"
1388
+ else:
1389
+ variable_gb = kv_cache_gb + activations_gb + optimizer_gb + gradients_gb
1390
+ variable_label = "KV + Activations + Optimizer + Gradients"
1391
+
1392
+ total_vram_gb = weights_gb + variable_gb
1393
+
1394
+ return total_vram_gb, weights_gb, variable_gb, variable_label, actual_precision
1395
+
1396
+ # =================================================================================================
1397
+ # Hardware Recommendation Engine
1398
+ # =================================================================================================
1399
+
1400
+ def recommend_hardware(
1401
+ required_vram: float,
1402
+ task: str,
1403
+ spec: ModelSpec,
1404
+ weights_gb: float,
1405
+ batch_size: int,
1406
+ pricing_tier: str,
1407
+ precision: str = "fp16",
1408
+ framework: str = "vllm",
1409
+ sample_count: int = 0,
1410
+ input_tokens: int = 0,
1411
+ output_tokens: int = 0,
1412
+ ft_method: Optional[str] = None,
1413
+ rank: Optional[int] = None,
1414
+ manufacturers: Optional[list[str]] = None,
1415
+ ) -> Tuple[Optional[str], str, str, Dict[str, Any]]:
1416
+ """
1417
+ Recommend GPU configurations based on VRAM requirements.
1418
+
1419
+ Args:
1420
+ required_vram: Required VRAM in GB
1421
+ task: Task type
1422
+ spec: Model specification
1423
+ weights_gb: Model weights in GB
1424
+ batch_size: Batch size
1425
+ pricing_tier: Pricing tier selection
1426
+ precision: Quantization/precision format
1427
+ framework: Inference framework
1428
+ sample_count: Number of samples (for time estimation)
1429
+ input_tokens: Input tokens per sample
1430
+ output_tokens: Output tokens per sample
1431
+ ft_method: Fine-tuning method (for training tasks)
1432
+ rank: LoRA rank (for LoRA/QLoRA training)
1433
+
1434
+ Returns:
1435
+ Tuple of (error_message, budget_rec, runner_up_rec, chart_data)
1436
+ """
1437
+ # Manufacturer filtering
1438
+ selected = manufacturers or []
1439
+ if not selected:
1440
+ selected = ["Nvidia"]
1441
+ filtered_gpus = [
1442
+ gpu for gpu in GPU_DATABASE
1443
+ if get_manufacturer(gpu.name) in selected
1444
+ ]
1445
+
1446
+ # VRAM filtering with 10% headroom
1447
+ valid_configs = [
1448
+ gpu for gpu in filtered_gpus
1449
+ if gpu.vram >= required_vram * 1.1
1450
+ ]
1451
+
1452
+ if not valid_configs:
1453
+ maxvram = max(g.vram for g in filtered_gpus) if filtered_gpus else 0
1454
+ errormsg = (
1455
+ f'<div class="error-box">'
1456
+ f"<h4>ERROR No suitable GPU configuration found</h4>"
1457
+ f"<p>Required VRAM <strong>{required_vram:.1f} GB</strong></p>"
1458
+ f"<p>The largest available configuration in the selected manufacturers "
1459
+ f"has {maxvram} GB VRAM.</p>"
1460
+ f"<p><em>Suggestions</em></p>"
1461
+ f"<ul>"
1462
+ f"<li>Reduce batch size (current {batch_size})</li>"
1463
+ f"<li>Use more aggressive quantization</li>"
1464
+ f"<li>Consider model sharding across multiple nodes</li>"
1465
+ f"<li>Try including more GPU manufacturers</li>"
1466
+ f"</ul>"
1467
+ f"</div>"
1468
+ )
1469
+ return errormsg, "", "", {}
1470
+
1471
+ # Sort by cost
1472
+ valid_by_cost = sorted(valid_configs, key=lambda x: x.get_price(pricing_tier))
1473
+
1474
+ # Build chart data for all valid configurations (limit to top 10 by cost for readability)
1475
+ chart_configs = valid_by_cost[:10]
1476
+ chart_data = {
1477
+ "names": [],
1478
+ "throughput": [],
1479
+ "cost": [],
1480
+ "vram_util": [],
1481
+ "cost_efficiency": [], # tokens per rupee
1482
+ }
1483
+
1484
+ for config in chart_configs:
1485
+ price = config.get_price(pricing_tier)
1486
+ vram_util = (required_vram / config.vram) * 100
1487
+ _, total_tps, _ = calculate_throughput(
1488
+ config, spec, task, batch_size, precision, framework, ft_method, rank
1489
+ )
1490
+
1491
+ chart_data["names"].append(config.name)
1492
+ chart_data["throughput"].append(total_tps)
1493
+ chart_data["cost"].append(price)
1494
+ chart_data["vram_util"].append(vram_util)
1495
+ # Cost efficiency: tokens per rupee per hour
1496
+ chart_data["cost_efficiency"].append(total_tps / price if price > 0 else 0)
1497
+
1498
+ # Generate recommendation cards
1499
+ def make_card(title: str, emoji: str, config: Optional[GPUConfig]) -> str:
1500
+ if config is None:
1501
+ return f"### {emoji} {title}\n*No configuration available*"
1502
+
1503
+ price = config.get_price(pricing_tier)
1504
+ vram_util = (required_vram / config.vram) * 100
1505
+
1506
+ tps_per_gpu, total_tps, throughput_desc = calculate_throughput(
1507
+ config, spec, task, batch_size, precision, framework, ft_method, rank
1508
+ )
1509
+
1510
+ time_estimate = ""
1511
+ if sample_count > 0 and (input_tokens > 0 or output_tokens > 0):
1512
+ total_tokens = sample_count * (input_tokens + output_tokens)
1513
+ time_str = format_time_estimate(total_tokens, total_tps)
1514
+ time_estimate = f"\n- **Time Estimate:** {time_str} for {sample_count:,} samples"
1515
+
1516
+ return f"""
1517
+ ### {emoji} {title}
1518
+ **{config.name}**
1519
+ - VRAM: {config.vram} GB ({config.vram_per_gpu:.0f} GB/GPU)
1520
+ - VRAM Utilization: {vram_util:.1f}%
1521
+ - Performance: {config.tflops:,.0f} TFLOPS
1522
+ - Bandwidth: {config.bandwidth:,.0f} GB/s
1523
+ - **Throughput: {total_tps:,.0f} tokens/sec**
1524
+ - Per GPU: {tps_per_gpu:,.0f} tok/s
1525
+ - {throughput_desc}{time_estimate}
1526
+ - **Price: β‚Ή{price:,.2f}/hour** from [IndiaAI price list](https://staging2.pmgatishakti.gov.in/IndiaAICompute/pricelist)
1527
+ - Daily: β‚Ή{price*24:,.2f} | Monthly: β‚Ή{price*730:,.2f}
1528
+ """
1529
+
1530
+ budget_rec = make_card("Best Budget", "πŸ₯‡", valid_by_cost[0] if valid_by_cost else None)
1531
+ runner_up_rec = make_card("Budget Runner-up", "πŸ₯ˆ", valid_by_cost[1] if len(valid_by_cost) > 1 else None)
1532
+
1533
+ return None, budget_rec, runner_up_rec, chart_data
1534
+
1535
+ # =================================================================================================
1536
+ # Get Manufacturer Function
1537
+ # =================================================================================================
1538
+
1539
+ def get_manufacturer(gpu_name: str) -> str:
1540
+ """Infer GPU manufacturer from config.name."""
1541
+ name = gpu_name.lower()
1542
+ if "nvidia" in name:
1543
+ return "Nvidia"
1544
+ if "amd" in name:
1545
+ return "AMD"
1546
+ if "intel" in name or "gaudi" in name:
1547
+ return "Intel"
1548
+ return "Other"
1549
+
1550
+ # =================================================================================================
1551
+ # Main Processing Function
1552
+ # =================================================================================================
1553
+
1554
+ def process_request(
1555
+ model_name: str,
1556
+ seq_len: float,
1557
+ quant: str,
1558
+ task: str,
1559
+ fw: str,
1560
+ ft_method: str,
1561
+ rank: float,
1562
+ batch: float,
1563
+ sample_count: float,
1564
+ input_tokens: float,
1565
+ output_tokens: float,
1566
+ dataset_tier: str,
1567
+ manufacturers: list[str],
1568
+ ) -> Tuple[str, str, str, Any, Any, Any, Any]:
1569
+ """
1570
+ Main request processing function.
1571
+
1572
+ Orchestrates model resolution, VRAM calculation, and hardware recommendation.
1573
+
1574
+ Returns:
1575
+ Tuple of (report, budget_rec, runner_up_rec, throughput_chart, cost_chart, vram_chart, efficiency_chart)
1576
+ """
1577
+ try:
1578
+ # Validate inputs
1579
+ validation_warnings = validate_inputs(seq_len, batch, rank, sample_count, input_tokens, output_tokens)
1580
+
1581
+ # Convert to integers
1582
+ seq_len = int(seq_len)
1583
+ batch = int(batch)
1584
+ rank = int(rank)
1585
+ sample_count = int(sample_count)
1586
+ input_tokens = int(input_tokens)
1587
+ output_tokens = int(output_tokens)
1588
+
1589
+ # Resolve model
1590
+ spec, source, logs = resolve_model(model_name, HF_TOKEN)
1591
+
1592
+ # Check if quantization is compatible with task
1593
+ if task == "Training" and quant in INFERENCE_ONLY_QUANT:
1594
+ quant = 'nf4'
1595
+ logs.append(f"WARNING: Auto-switched to 'nf4' ({quant} is inference-only)")
1596
+
1597
+ # Set ft_method to None for inference
1598
+ if task == "Inference":
1599
+ ft_method = None
1600
+
1601
+ # Add explanation for training precision
1602
+ if task == "Training" and ft_method:
1603
+ if ft_method == "QLoRA":
1604
+ logs.append(f"INFO: QLoRA - Base model stays quantized ({quant}), adapters in fp16")
1605
+ elif ft_method == "LoRA":
1606
+ if quant in ['nf4', '4bit', 'int4', 'int8']:
1607
+ logs.append(f"INFO: LoRA - Base model upgraded to bf16 (LoRA requires full precision)")
1608
+ else:
1609
+ logs.append(f"INFO: LoRA - Base model in {quant}, adapters in fp16")
1610
+ elif ft_method == "Full Fine-Tuning":
1611
+ if quant in ['nf4', '4bit', 'int4', 'int8', 'awq', 'gptq']:
1612
+ logs.append(f"INFO: Full FT - Base model upgraded to bf16 (training requires full precision)")
1613
+ else:
1614
+ logs.append(f"INFO: Full Fine-Tuning - Training all parameters at {quant}")
1615
+
1616
+ # Calculate VRAM
1617
+ total_vram, weights_gb, variable_gb, variable_label, actual_precision = calculate_vram(
1618
+ spec, quant, batch, seq_len, task, fw, ft_method, rank
1619
+ )
1620
+
1621
+ # Check token length warning
1622
+ tokens_per_sample = input_tokens + output_tokens
1623
+ if tokens_per_sample > seq_len:
1624
+ logs.append(f"WARNING: Sample length ({tokens_per_sample}) exceeds context length ({seq_len})")
1625
+
1626
+ # Generate hardware recommendations
1627
+ error, budget_rec, runner_up_rec, chart_data = recommend_hardware(
1628
+ total_vram, task, spec, weights_gb, batch, dataset_tier, actual_precision, fw,
1629
+ sample_count, input_tokens, output_tokens, ft_method, rank, manufacturers=manufacturers,
1630
+ )
1631
+
1632
+ if error:
1633
+ return error, "", "", None, None, None, None
1634
+
1635
+ # Generate report
1636
+ report = f"""
1637
+ ### πŸ“Š Analysis Report
1638
+
1639
+ **Model Information:**
1640
+ - Source: {source}
1641
+ - Parameters: **{spec.params_bn:.2f}B**
1642
+ - Layers: {spec.layers} | Heads: {spec.heads} | KV Heads: {spec.kv_heads}
1643
+ - Context Length: {spec.context:,} tokens
1644
+
1645
+ **VRAM Breakdown:**
1646
+ - Model Weights: **{weights_gb:.1f} GB** ({quant})
1647
+ - {variable_label}: **{variable_gb:.1f} GB**
1648
+ - **Total Required: {total_vram:.1f} GB**
1649
+
1650
+ **Configuration:**
1651
+ - Task: {task}
1652
+ - Framework: {fw}
1653
+ - Batch Size: {batch}
1654
+ - Sequence Length: {seq_len:,}
1655
+ {f"- Fine-tuning: {ft_method} (rank={rank})" if ft_method in ["LoRA", "QLoRA"] else f"- Fine-tuning: {ft_method}" if ft_method else ""}
1656
+
1657
+ ---
1658
+ {"<br>".join(f"*{log}*" for log in logs)}
1659
+ {"<br>".join(f'<div class="warning-box">{w}</div>' for w in validation_warnings) if validation_warnings else ""}
1660
+
1661
+ *ℹ️ NOTE: VRAM estimates include 10% safety buffer. Actual usage may vary Β±15-20% based on framework optimizations, model architecture details, and runtime conditions. Throughput estimates are based on empirical benchmarks and may vary in production.*
1662
+ """
1663
+
1664
+ # Create bar charts using Plotly
1665
+ throughput_chart = None
1666
+ cost_chart = None
1667
+ vram_chart = None
1668
+ efficiency_chart = None
1669
+
1670
+ if chart_data and chart_data.get("names"):
1671
+ import plotly.graph_objects as go
1672
+
1673
+ # Throughput chart
1674
+ throughput_chart = go.Figure(data=[
1675
+ go.Bar(
1676
+ x=chart_data["names"],
1677
+ y=chart_data["throughput"],
1678
+ marker_color='#22c55e'
1679
+ )
1680
+ ])
1681
+ throughput_chart.update_layout(
1682
+ title="Throughput Comparison",
1683
+ xaxis_title="GPU Configuration",
1684
+ yaxis_title="Tokens/sec",
1685
+ height=350,
1686
+ xaxis_tickangle=-45,
1687
+ margin=dict(b=120)
1688
+ )
1689
+
1690
+ # Cost chart
1691
+ cost_chart = go.Figure(data=[
1692
+ go.Bar(
1693
+ x=chart_data["names"],
1694
+ y=chart_data["cost"],
1695
+ marker_color='#3b82f6'
1696
+ )
1697
+ ])
1698
+ cost_chart.update_layout(
1699
+ title="Cost Comparison",
1700
+ xaxis_title="GPU Configuration",
1701
+ yaxis_title="β‚Ή/hour",
1702
+ height=350,
1703
+ xaxis_tickangle=-45,
1704
+ margin=dict(b=120)
1705
+ )
1706
+
1707
+ # VRAM utilization chart
1708
+ vram_chart = go.Figure(data=[
1709
+ go.Bar(
1710
+ x=chart_data["names"],
1711
+ y=chart_data["vram_util"],
1712
+ marker_color='#a855f7'
1713
+ )
1714
+ ])
1715
+ vram_chart.update_layout(
1716
+ title="VRAM Utilization Comparison",
1717
+ xaxis_title="GPU Configuration",
1718
+ yaxis_title="Utilization %",
1719
+ height=350,
1720
+ xaxis_tickangle=-45,
1721
+ margin=dict(b=120)
1722
+ )
1723
+
1724
+ # Cost efficiency chart
1725
+ efficiency_chart = go.Figure(data=[
1726
+ go.Bar(
1727
+ x=chart_data["names"],
1728
+ y=chart_data["cost_efficiency"],
1729
+ marker_color='#f59e0b'
1730
+ )
1731
+ ])
1732
+ efficiency_chart.update_layout(
1733
+ title="Cost Efficiency Comparison",
1734
+ xaxis_title="GPU Configuration",
1735
+ yaxis_title="Throughput / Cost",
1736
+ height=350,
1737
+ xaxis_tickangle=-45,
1738
+ margin=dict(b=120)
1739
+ )
1740
+
1741
+ return report, budget_rec, runner_up_rec, throughput_chart, cost_chart, vram_chart, efficiency_chart
1742
+
1743
+ except Exception as e:
1744
+ error_msg = f"""
1745
+ <div class="error-box">
1746
+ <h4>ERROR: Error Processing Request</h4>
1747
+ <p>{str(e)}</p>
1748
+ </div>
1749
+ """
1750
+ return error_msg, "", "", None, None, None, None
1751
+
1752
+ # =================================================================================================
1753
+ # UI Event Handlers
1754
+ # =================================================================================================
1755
+
1756
+ def update_ui_on_task(task_val: str):
1757
+ """Update UI elements when task changes."""
1758
+ if task_val == "Training":
1759
+ return [
1760
+ gr.update(choices=["nf4", "int4", "bf16", "fp16", "int8"], value="nf4"),
1761
+ gr.update(choices=["huggingface"], value="huggingface"),
1762
+ gr.update(visible=True)
1763
+ ]
1764
+ else:
1765
+ return [
1766
+ gr.update(choices=["nf4", "int4", "bf16", "fp16", "int8", "awq", "gptq"], value="nf4"),
1767
+ gr.update(choices=["vllm"], value="vllm"),
1768
+ gr.update(visible=False)
1769
+ ]
1770
+
1771
+ def update_rank_visibility(ft_val: str):
1772
+ """Update rank input visibility based on fine-tuning method."""
1773
+ if ft_val in ["LoRA", "QLoRA"]:
1774
+ return gr.update(visible=True)
1775
+ return gr.update(visible=False)
1776
+
1777
+ # =================================================================================================
1778
+ # Gradio Interface
1779
+ # =================================================================================================
1780
+
1781
+ def create_interface() -> gr.Blocks:
1782
+ """Create and configure the Gradio interface."""
1783
+
1784
+ # Try to create with theme, fallback to basic if not supported
1785
+ try:
1786
+ demo = gr.Blocks(title="IndiaAI GPU Infrastructure Recommender")
1787
+ except Exception as e:
1788
+ print(f"Note: Using basic Gradio configuration: {e}")
1789
+ demo = gr.Blocks()
1790
+
1791
+ with demo:
1792
+ # Inject custom CSS
1793
+ gr.HTML(CUSTOM_CSS)
1794
+
1795
+ # Header
1796
+ gr.Markdown("""
1797
+ # IndiaAI GPU Infrastructure Recommender
1798
+
1799
+ Calculate VRAM requirements and get optimal GPU recommendations for your AI workloads using [IndiaAI's price list](https://staging2.pmgatishakti.gov.in/IndiaAICompute/pricelist).
1800
+ """)
1801
+
1802
+ with gr.Row():
1803
+ # Left column: Inputs
1804
+ with gr.Column(scale=0.30):
1805
+ gr.Markdown("### πŸ”§ Configuration")
1806
+
1807
+ model_name = gr.Dropdown(
1808
+ choices=MODEL_CHOICES,
1809
+ label="Model Name or Size",
1810
+ value="mistralai/Mistral-7B-Instruct-v0.3",
1811
+ allow_custom_value=True,
1812
+ info="Select a model or enter custom (e.g., '7B', 'meta-llama/...')"
1813
+ )
1814
+
1815
+ with gr.Row():
1816
+ task = gr.Radio(
1817
+ ["Inference", "Training"],
1818
+ label="Task",
1819
+ value="Inference",
1820
+ info="Select your use case"
1821
+ )
1822
+ fw = gr.Dropdown(
1823
+ ["vllm"], #, "huggingface"],
1824
+ label="Framework",
1825
+ value="vllm",
1826
+ info="Framework for the task"
1827
+ )
1828
+ quant = gr.Dropdown(
1829
+ ["nf4", "int4", "bf16", "fp16", "int8", "awq", "gptq"],
1830
+ label="Quantization",
1831
+ value="nf4",
1832
+ info="Precision format"
1833
+ )
1834
+
1835
+ with gr.Row(visible=False) as training_row:
1836
+ ft_method = gr.Dropdown(
1837
+ ["QLoRA", "LoRA", "Full Fine-Tuning"],
1838
+ label="Fine-tuning Method",
1839
+ value="QLoRA",
1840
+ info="Training strategy"
1841
+ )
1842
+ rank = gr.Number(
1843
+ label="LoRA Rank (r)",
1844
+ value=16,
1845
+ visible=True,
1846
+ info="Adapter rank for LoRA/QLoRA"
1847
+ )
1848
+
1849
+ gr.Markdown("### πŸ“ Workload Parameters")
1850
+
1851
+ with gr.Row():
1852
+ seq_len = gr.Number(
1853
+ label="Max Context Length",
1854
+ value=1024,
1855
+ info="Maximum sequence length (buffer)"
1856
+ )
1857
+ batch = gr.Number(
1858
+ label="Batch Size",
1859
+ value=16,
1860
+ info="Number of samples processed together"
1861
+ )
1862
+
1863
+ with gr.Row():
1864
+ sample_count = gr.Number(
1865
+ label="Samples",
1866
+ value=10000,
1867
+ scale=1,
1868
+ info="Number of samples"
1869
+ )
1870
+ input_tokens = gr.Number(
1871
+ label="Input Tokens",
1872
+ value=300,
1873
+ scale=1,
1874
+ info="Avg input length"
1875
+ )
1876
+ output_tokens = gr.Number(
1877
+ label="Output Tokens",
1878
+ value=100,
1879
+ scale=1,
1880
+ info="Avg output length"
1881
+ )
1882
+
1883
+ gr.Markdown("### πŸ’° Cost Estimation")
1884
+
1885
+ with gr.Row():
1886
+ dataset_tier = gr.Dropdown(
1887
+ choices=["On Demand", "1 Month Reserved", "6 Month Reserved", "12 Month Reserved"],
1888
+ label="Pricing Tier",
1889
+ value="On Demand",
1890
+ scale=2,
1891
+ info="Select pricing model",
1892
+ )
1893
+ manufacturers = gr.CheckboxGroup(
1894
+ choices=["Nvidia", "AMD", "Intel"],
1895
+ label="GPU Manufacturers",
1896
+ value=["Nvidia", "AMD", "Intel"], # default: all
1897
+ info="Filter recommendations by GPU manufacturer",
1898
+ )
1899
+
1900
+ btn = gr.Button("Calculate Requirements", variant="primary")
1901
+
1902
+ # Right column: Results
1903
+ with gr.Column(scale=1):
1904
+ gr.HTML("<h3 style='text-align: center; margin-top: 0px;'>Recommendations</h3>")
1905
+
1906
+ with gr.Row():
1907
+ with gr.Column(elem_classes=["budget-box"]):
1908
+ rec_out_1 = gr.Markdown()
1909
+ with gr.Column(elem_classes=["runner-box"]):
1910
+ rec_out_2 = gr.Markdown()
1911
+
1912
+ # Bar chart comparison section with tabs inside Accordion
1913
+ with gr.Accordion("πŸ“Š GPU Comparison (Top 10 by Cost)", open=False):
1914
+ with gr.Tabs():
1915
+ with gr.Tab("Throughput"):
1916
+ throughput_plot = gr.Plot(label="Throughput Comparison")
1917
+ with gr.Tab("Cost"):
1918
+ cost_plot = gr.Plot(label="Cost Comparison")
1919
+ with gr.Tab("VRAM Utilization"):
1920
+ vram_plot = gr.Plot(label="VRAM Utilization Comparison")
1921
+ with gr.Tab("Cost Efficiency"):
1922
+ efficiency_plot = gr.Plot(label="Cost Efficiency Comparison")
1923
+
1924
+ with gr.Accordion("πŸ“ˆ Details", open=False):
1925
+ report_out = gr.Markdown()
1926
+
1927
+ # Footer
1928
+ gr.Markdown("""
1929
+ ---
1930
+ πŸ’‘ **Tips:**
1931
+ - Start with smaller batch sizes for testing
1932
+ - QLoRA is most memory-efficient for training
1933
+ - Consider reserved instances for long-term workloads
1934
+ - VRAM estimates include 10% safety buffer
1935
+ - Higher "Cost Efficiency" means more tokens processed per rupee spent
1936
+ """)
1937
+
1938
+ # Event handlers
1939
+ task.change(
1940
+ fn=update_ui_on_task,
1941
+ inputs=task,
1942
+ outputs=[quant, fw, training_row]
1943
+ )
1944
+
1945
+ ft_method.change(
1946
+ fn=update_rank_visibility,
1947
+ inputs=ft_method,
1948
+ outputs=rank
1949
+ )
1950
+
1951
+ btn.click(
1952
+ fn=process_request,
1953
+ inputs=[
1954
+ model_name, seq_len, quant, task, fw,
1955
+ ft_method, rank, batch, sample_count,
1956
+ input_tokens, output_tokens, dataset_tier, manufacturers,
1957
+ ],
1958
+ outputs=[report_out, rec_out_1, rec_out_2, throughput_plot, cost_plot, vram_plot, efficiency_plot]
1959
+ )
1960
+
1961
+ return demo
1962
+
1963
+ # =================================================================================================
1964
+ # Main Entry Point
1965
+ # =================================================================================================
1966
+
1967
+ if __name__ == "__main__":
1968
+ demo = create_interface()
1969
+ demo.launch(share=False)
requirements.txt ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ gradio>=4.0.0
2
+ transformers>=4.30.0
3
+ huggingface-hub>=0.16.0
4
+ plotly>=5.0.0