"""Pure contracts shared by model loading and runtime selection paths.""" from __future__ import annotations from collections.abc import Collection from dataclasses import dataclass from typing import Literal VALID_MODEL_TASKS = ("causal_lm", "classification") VALID_QUANTIZATIONS = (None, "4bit", "8bit") VALID_DTYPES = ("float32", "float16", "bfloat16") PACKED_QUANTIZED_MODULE_TYPES = frozenset( {"QuantLinear", "WQLinear", "WQLinear_GEMM", "WQLinear_GEMV"}, ) QUANTIZED_PARAMETER_TYPES = frozenset( {"Params4bit", "Int8Params", *PACKED_QUANTIZED_MODULE_TYPES}, ) SHARED_KV_PROJECTION_NAMES = frozenset({"k_proj", "v_proj", "k_norm"}) @dataclass(frozen=True) class ModelLoadPolicy: """Provider-independent decisions that control model placement and loading.""" quantization_backend: Literal["none", "native", "bitsandbytes"] include_torch_dtype: bool use_device_map_auto: bool move_to_resolved_device: bool def validate_model_load_request( model_name: object, task: str, quantization: str | None, dtype: str, *, valid_tasks: Collection[str] = VALID_MODEL_TASKS, ) -> None: """Reject malformed loader requests before any provider access.""" if not isinstance(model_name, str) or not model_name.strip(): raise ValueError("model_name must be a non-empty HuggingFace identifier or local path") if task not in valid_tasks: raise ValueError(f"Unknown task {task!r}. Choose from {list(valid_tasks)}") if quantization not in VALID_QUANTIZATIONS: raise ValueError( "Unknown quantization {!r}. Choose None, '4bit', or '8bit'".format(quantization), ) if dtype not in VALID_DTYPES: raise ValueError(f"Unknown dtype {dtype!r}. Choose from {list(VALID_DTYPES)}") def resolve_model_load_policy( *, device: str, resolved_device: str, dtype: str, quantization: str | None, has_native_quantization: bool, device_map_auto_supported: bool, bitsandbytes_supported: bool, ) -> ModelLoadPolicy: """Resolve dtype, quantization, device-map, and post-load move precedence.""" if quantization not in VALID_QUANTIZATIONS: raise ValueError( "Unknown quantization {!r}. Choose None, '4bit', or '8bit'".format(quantization), ) if has_native_quantization: return ModelLoadPolicy( quantization_backend="native", include_torch_dtype=False, use_device_map_auto=True, move_to_resolved_device=False, ) if quantization in ("4bit", "8bit"): if not bitsandbytes_supported: raise RuntimeError( f"Quantization '{quantization}' requires an available NVIDIA CUDA device; " f"resolved device was '{resolved_device}'. Remove --quantization to load in {dtype}.", ) return ModelLoadPolicy( quantization_backend="bitsandbytes", include_torch_dtype=True, use_device_map_auto=True, move_to_resolved_device=False, ) use_device_map_auto = device == "auto" and device_map_auto_supported return ModelLoadPolicy( quantization_backend="none", include_torch_dtype=True, use_device_map_auto=use_device_map_auto, move_to_resolved_device=not use_device_map_auto, ) def attention_projection_names( *, projection_target: str, layer_index: int, num_layers: int, num_kv_shared_layers: int, output_names: Collection[str], input_names: Collection[str], ) -> tuple[str, ...]: """Choose attention weights while projecting shared KV storage exactly once.""" outputs = tuple(output_names) if projection_target == "output": return outputs all_names = outputs + tuple(input_names) if num_kv_shared_layers <= 0: return all_names if not 0 <= layer_index < num_layers: raise ValueError("shared-KV projection requires a valid layer index and layer count") if num_kv_shared_layers > num_layers: raise ValueError("num_kv_shared_layers cannot exceed the model layer count") owner_index = num_layers - num_kv_shared_layers if layer_index <= owner_index: return all_names return tuple(name for name in all_names if name not in SHARED_KV_PROJECTION_NAMES) def is_quantized_parameter(*, class_name: str, has_quant_state: bool) -> bool: """Return whether a parameter carries a supported packed-quantization marker.""" return has_quant_state or class_name in QUANTIZED_PARAMETER_TYPES def classify_weight_storage( *, module_class_name: str, parameter_class_name: str, has_quant_state: bool, data_is_floating_point: bool, ) -> Literal["packed_module", "quantized_parameter", "integer", "float"]: """Classify a projection weight so callers select a safe read/write path.""" if module_class_name in PACKED_QUANTIZED_MODULE_TYPES: return "packed_module" if is_quantized_parameter( class_name=parameter_class_name, has_quant_state=has_quant_state, ): return "quantized_parameter" if not data_is_floating_point: return "integer" return "float" def effective_model_memory_gb(estimate_gb: float, quantization: str | None) -> float: """Adjust a full-precision weight estimate for runtime quantization.""" factor = {"4bit": 4, "8bit": 2}.get(quantization, 1) return estimate_gb / factor def quantized_model_fits_gpu( estimate_gb: float, quantization: str | None, available_gpu_gb: float, ) -> bool: """Return whether a quantized estimate fits with 30 percent headroom.""" effective_gb = effective_model_memory_gb(estimate_gb, quantization) return ( quantization in ("4bit", "8bit") and effective_gb > 0 and effective_gb < available_gpu_gb * 0.7 ) def should_snapshot_model( *, skip_snapshot: bool | None, initial_gpu_free_gb: float, remaining_gpu_free_gb: float, has_native_quantization: bool, estimate_gb: float, quantization: str | None, ) -> bool: """Decide whether a restorable state snapshot fits the memory policy.""" if skip_snapshot is True: return False if skip_snapshot is False: return True if initial_gpu_free_gb > 0 and has_native_quantization: return remaining_gpu_free_gb >= initial_gpu_free_gb * 0.4 if initial_gpu_free_gb > 0: effective_gb = effective_model_memory_gb(estimate_gb, quantization) return effective_gb <= initial_gpu_free_gb * 0.5 return True def classify_architecture_size( *, is_moe: bool, total_params_b: float, num_experts: int, model_name: str, large_moe_name_patterns: Collection[str], ) -> Literal["dense", "small_moe", "large_moe"]: """Classify dense and MoE scale using the documented precedence rules.""" if not is_moe: return "dense" if total_params_b > 0: is_small = total_params_b < 100 elif num_experts > 0: is_small = num_experts <= 16 else: name_lower = model_name.lower() is_small = not any(pattern.lower() in name_lower for pattern in large_moe_name_patterns) return "small_moe" if is_small else "large_moe" def supports_bfloat16_target( device: str, *, cuda_available: bool, cuda_major: int | None, torch_version: str, ) -> bool: """Evaluate the bfloat16 capability contract from deterministic inputs.""" if device.startswith("cuda"): return cuda_available and cuda_major is not None and cuda_major >= 8 if device == "mps": major, minor = (int(value) for value in torch_version.split(".")[:2]) return (major, minor) >= (2, 3) return True