splitbit-llm / splitbit_llm /train /auto_size.py
hermescures1's picture
Upload folder using huggingface_hub
0e3d4b8 verified
Raw
History Blame Contribute Delete
7.39 kB
"""Auto-size adjustment for SplitBit LLM.
Detects hardware (reuses config.py hardware detection).
Picks optimal model config per tier.
Adjusts batch size and training hyperparams accordingly.
Can upgrade/downgrade model size at runtime (transfer weights).
"""
from __future__ import annotations
import logging
from typing import Any
from ..config import (
HardwareTier,
ModelConfig,
QuantConfig,
StorageConfig,
detect_hardware,
get_model_config,
)
logger = logging.getLogger(__name__)
# Training hyperparameters per tier
TIER_TRAINING_CONFIG: dict[HardwareTier, dict[str, Any]] = {
HardwareTier.MOBILE: {
"batch_size": 1,
"seq_len": 128,
"lr": 1e-3,
"epochs": 20,
"warmup_steps": 50,
},
HardwareTier.MINIMAL: {
"batch_size": 2,
"seq_len": 256,
"lr": 3e-4,
"epochs": 15,
"warmup_steps": 100,
},
HardwareTier.LIGHT: {
"batch_size": 4,
"seq_len": 256,
"lr": 3e-4,
"epochs": 10,
"warmup_steps": 200,
},
HardwareTier.STANDARD: {
"batch_size": 8,
"seq_len": 512,
"lr": 3e-4,
"epochs": 10,
"warmup_steps": 500,
},
HardwareTier.FULL: {
"batch_size": 16,
"seq_len": 1024,
"lr": 2e-4,
"epochs": 8,
"warmup_steps": 1000,
},
HardwareTier.MAXIMUM: {
"batch_size": 32,
"seq_len": 2048,
"lr": 1e-4,
"epochs": 5,
"warmup_steps": 2000,
},
HardwareTier.DATACENTER: {
"batch_size": 64,
"seq_len": 4096,
"lr": 5e-5,
"epochs": 3,
"warmup_steps": 5000,
},
HardwareTier.SUPERCOMPUTER: {
"batch_size": 128,
"seq_len": 8192,
"lr": 2e-5,
"epochs": 2,
"warmup_steps": 10000,
},
}
# Inference parameters per tier
TIER_INFERENCE_CONFIG: dict[HardwareTier, dict[str, Any]] = {
HardwareTier.MOBILE: {
"max_tokens": 32,
"temperature": 0.5,
"top_k": 20,
"use_cache": True,
},
HardwareTier.MINIMAL: {
"max_tokens": 64,
"temperature": 0.7,
"top_k": 40,
"use_cache": True,
},
HardwareTier.LIGHT: {
"max_tokens": 128,
"temperature": 0.7,
"top_k": 40,
"use_cache": True,
},
HardwareTier.STANDARD: {
"max_tokens": 256,
"temperature": 0.7,
"top_k": 50,
"use_cache": True,
},
HardwareTier.FULL: {
"max_tokens": 512,
"temperature": 0.7,
"top_k": 50,
"use_cache": True,
},
HardwareTier.MAXIMUM: {
"max_tokens": 1024,
"temperature": 0.7,
"top_k": 100,
"use_cache": True,
},
HardwareTier.DATACENTER: {
"max_tokens": 2048,
"temperature": 0.7,
"top_k": 100,
"use_cache": True,
},
HardwareTier.SUPERCOMPUTER: {
"max_tokens": 4096,
"temperature": 0.7,
"top_k": 200,
"use_cache": True,
},
}
# Voice-optimized inference parameters per tier
TIER_VOICE_CONFIG: dict[HardwareTier, dict[str, Any]] = {
HardwareTier.MOBILE: {
"max_tokens": 16,
"temperature": 0.3,
"top_k": 10,
"use_cache": True,
},
HardwareTier.MINIMAL: {
"max_tokens": 32,
"temperature": 0.5,
"top_k": 20,
"use_cache": True,
},
HardwareTier.LIGHT: {
"max_tokens": 48,
"temperature": 0.5,
"top_k": 30,
"use_cache": True,
},
HardwareTier.STANDARD: {
"max_tokens": 64,
"temperature": 0.5,
"top_k": 40,
"use_cache": True,
},
HardwareTier.FULL: {
"max_tokens": 96,
"temperature": 0.5,
"top_k": 50,
"use_cache": True,
},
HardwareTier.MAXIMUM: {
"max_tokens": 128,
"temperature": 0.5,
"top_k": 50,
"use_cache": True,
},
HardwareTier.DATACENTER: {
"max_tokens": 256,
"temperature": 0.5,
"top_k": 100,
"use_cache": True,
},
HardwareTier.SUPERCOMPUTER: {
"max_tokens": 512,
"temperature": 0.5,
"top_k": 200,
"use_cache": True,
},
}
class AutoSizer:
"""Hardware-aware model sizing and hyperparameter adjustment.
Detects hardware tier and provides optimal configs for:
- Model architecture (layers, heads, dim, vocab)
- Weight quantization format
- Training hyperparameters (batch size, lr, epochs)
- Inference parameters (max tokens, temperature)
- Voice-optimized inference (shorter, faster responses)
- Skill storage limits
"""
def __init__(self, tier: HardwareTier | None = None) -> None:
self.tier = tier or detect_hardware()
self.model_config = get_model_config(self.tier)
self.quant_config = QuantConfig.for_tier(self.tier)
self.storage_config = StorageConfig.for_tier(self.tier)
self.training_config = TIER_TRAINING_CONFIG.get(self.tier, TIER_TRAINING_CONFIG[HardwareTier.MINIMAL])
self.inference_config = TIER_INFERENCE_CONFIG.get(self.tier, TIER_INFERENCE_CONFIG[HardwareTier.MINIMAL])
self.voice_config = TIER_VOICE_CONFIG.get(self.tier, TIER_VOICE_CONFIG[HardwareTier.MINIMAL])
logger.info(
"AutoSizer: tier=%s, model=%dL/%dH/d%d, quant=%s (%.1f bpw), "
"batch=%d, lr=%.1e, max_tokens=%d, voice_tokens=%d",
self.tier.value,
self.model_config.n_layers, self.model_config.n_heads, self.model_config.d_model,
self.quant_config.format, self.quant_config.bpw,
self.training_config["batch_size"], self.training_config["lr"],
self.inference_config["max_tokens"], self.voice_config["max_tokens"],
)
def get_model_config(self) -> ModelConfig:
return self.model_config
def get_quant_config(self) -> QuantConfig:
return self.quant_config
def get_training_params(self) -> dict[str, Any]:
return self.training_config.copy()
def get_inference_params(self) -> dict[str, Any]:
return self.inference_config.copy()
def get_voice_params(self) -> dict[str, Any]:
return self.voice_config.copy()
def get_storage_config(self) -> StorageConfig:
return self.storage_config
def get_all_stats(self) -> dict[str, Any]:
return {
"tier": self.tier.value,
"model": {
"n_layers": self.model_config.n_layers,
"n_heads": self.model_config.n_heads,
"d_model": self.model_config.d_model,
"d_ff": self.model_config.d_ff,
"vocab_size": self.model_config.vocab_size,
"max_seq_len": self.model_config.max_seq_len,
"param_count": self.model_config.param_count,
},
"quant": {
"format": self.quant_config.format,
"bpw": self.quant_config.bpw,
},
"training": self.training_config,
"inference": self.inference_config,
"voice": self.voice_config,
"storage": {
"skill_storage_mb": self.storage_config.skill_storage_mb,
"max_skills": self.storage_config.max_skills,
},
}