Setting the file. One moment. Estimate Cost · Huggingface LLM Trainer · huggingface/skills · Skills Docs14.10
Unsloth
scripts/estimate_cost.py
scripts/estimate_cost.py
Python·150 lines·5 KB
17
18# Hardware costs per hour (approximate)
19HARDWARE_COSTS = {
20 "t4-small": 0.75,
21 "t4-medium": 1.50,
22 "l4x1": 2.50,
23 "a10g-small": 3.50,
24 "a10g-large": 5.00,
25 "a10g-largex2": 10.00,
26 "a10g-largex4": 20.00,
27 "a100-large": 10.00,
28}
29
30# Model sizes in billions of parameters
31MODEL_SIZES = {
32 "0.5B": 0.5,
33 "1.5B": 1.5,
34 "3B": 3,
35 "7B": 7,
36 "13B": 13,
37}
38
39def estimate_training_time(model_params, dataset_size, epochs, hardware):
40 """Estimate training time in hours."""
41 # Rough estimates based on empirical observations
42 # These are approximations and actual times will vary
43
44 base_time_per_1k_examples = 0.1 # hours for 1B model on a10g-large
45
46 # Adjust for model size
47 time = base_time_per_1k_examples * model_params * (dataset_size / 1000) * epochs
48
49 # Adjust for hardware (relative to a10g-large baseline)
50 hardware_multipliers = {
51 "t4-small": 2.0,
52 "t4-medium": 1.5,
53 "l4x1": 1.2,
54 "a10g-small": 1.3,
55 "a10g-large": 1.0,
56 "a10g-largex2": 0.6,
57 "a10g-largex4": 0.4,
58 "a100-large": 0.7,
59 }
60
61 multiplier = hardware_multipliers.get(hardware, 1.0)
62 time *= multiplier
63
64 return time
65
66def parse_args():
67 parser = argparse.ArgumentParser(description="Estimate training cost for TRL jobs")
68 parser.add_argument("--model", required=True, help="Model name or size (e.g., 'Qwen/Qwen2.5-0.5B' or '0.5B')")
69 parser.add_argument("--dataset", required=True, help="Dataset name")
70 parser.add_argument("--hardware", required=True, choices=HARDWARE_COSTS.keys(), help="Hardware flavor")
71 parser.add_argument("--dataset-size", type=int, help="Override dataset size (number of examples)")
72 parser.add_argument("--epochs", type=int, default=3, help="Number of training epochs")
73 return parser.parse_args()
74
75def extract_model_size(model_name):
76 """Extract model size from name or return parsed value."""
77 for size_str, size_val in MODEL_SIZES.items():
78 if size_str in model_name:
79 return size_val
80
81 # Try to parse directly
82 try:
83 if "B" in model_name:
84 return float(model_name.replace("B", ""))
85 except:
86 pass
87
88 return 1.0 # Default to 1B if can't determine
89
90def main():
91 args = parse_args()
92
93 # Extract model parameters
94 model_params = extract_model_size(args.model)
95 print(f"📊 Model: {args.model} (~{model_params}B parameters)")
96
97 # Estimate dataset size (would need to load to get real size)
98 if args.dataset_size:
99 dataset_size = args.dataset_size
100 else:
101 # Common dataset sizes (approximations)
102 dataset_sizes = {
103 "trl-lib/Capybara": 16000,
104 "Anthropic/hh-rlhf": 160000,
105 }
106 dataset_size = dataset_sizes.get(args.dataset, 10000)
107
108 print(f"📦 Dataset: {args.dataset} (~{dataset_size} examples)")
109 print(f"🔄 Epochs: {args.epochs}")
110 print(f"💻 Hardware: {args.hardware}")
111 print()
112
113 # Estimate training time
114 estimated_hours = estimate_training_time(model_params, dataset_size, args.epochs, args.hardware)
115 estimated_cost = estimated_hours * HARDWARE_COSTS[args.hardware]
116
117 # Recommend timeout with buffer
118 recommended_timeout_hours = estimated_hours * 1.3 # 30% buffer
119
120 print(f"⏱️ Estimated training time: {estimated_hours:.1f} hours")
121 print(f"💰 Estimated cost: ${estimated_cost:.2f}")
122 print(f"⏰ Recommended timeout: {recommended_timeout_hours:.1f}h (with 30% buffer)")
123 print()
124
125 # Warnings and recommendations
126 if estimated_hours > 4:
127 print("⚠️ Long training time - consider:")
128 print(" - Using faster hardware")
129 print(" - Reducing epochs")
130 print(" - Using a smaller dataset subset for testing")
131
132 if model_params >= 7 and args.hardware not in ["a10g-largex2", "a10g-largex4", "a100-large"]:
133 print("⚠️ Large model - consider using:")
134 print(" - Larger GPU (a100-large)")
135 print(" - Multi-GPU setup (a10g-largex2 or a10g-largex4)")
136 print(" - LoRA/PEFT for memory efficiency")
137
138 print()
139 print("📋 Example job configuration:")
140 print(f"""
141hf_jobs("uv", {{
142 "script": "your_training_script.py",
143 "flavor": "{args.hardware}",
144 "timeout": "{recommended_timeout_hours:.0f}h",
145 "secrets": {{"HF_TOKEN": "$HF_TOKEN"}}
146}})
147""")
148
149if __name__ == "__main__":
150 main()