"""Integer lower-bound LoRA memory accounting, not a runtime peak predictor.""" from dataclasses import asdict, dataclass, fields import hashlib import json import math import sys def text(value): if type(value) is not str or not value.strip() or len(value) > 4096: raise ValueError("nonempty bounded string required") def count(value): if type(value) is not int or not 1 <= value <= 1_000_000: raise ValueError("bounded positive integer required") def score(value): if type(value) not in (float, int) or not math.isfinite(value) or not 0 <= value <= 1: raise ValueError("finite score in [0,1], not bool") if value != 0 and abs(value) < sys.float_info.min: raise ValueError("subnormal rejected") def digest(value): return hashlib.sha256(json.dumps(value, sort_keys=True, separators=(",", ":"), allow_nan=False).encode()).hexdigest() @dataclass(frozen=True) class Target: target_id: str input_dim: int output_dim: int rank: int def __post_init__(self): text(self.target_id) for item in (self.input_dim, self.output_dim, self.rank): count(item) if self.rank > min(self.input_dim, self.output_dim): raise ValueError("rank exceeds matrix dimensions") @dataclass(frozen=True) class BudgetContract: scope: str model: str data: str policy: str version: str provenance: str eval_digest: str eval_scope: str eval_model: str eval_data: str eval_config_digest: str eval_cases: int min_eval_cases: int held_out: bool quality: float required_quality: float base_parameters: int base_dtype: str adapter_dtype: str optimizer: str adapter_ids: tuple[str, ...] target_ids: tuple[str, ...] shard_count: int targets: tuple[Target, ...] def __post_init__(self): for name in ("scope", "model", "data", "policy", "version", "provenance", "eval_scope", "eval_model", "eval_data", "base_dtype", "adapter_dtype", "optimizer"): text(getattr(self, name)) for name in ("eval_digest", "eval_config_digest"): value = getattr(self, name) if type(value) is not str or len(value) != 64 or any(c not in "0123456789abcdef" for c in value): raise ValueError("content SHA-256 required") for value in (self.eval_cases, self.min_eval_cases, self.shard_count): count(value) if type(self.base_parameters) is not int or not 1 <= self.base_parameters <= 10**15: raise ValueError("bounded base parameter count required") if self.base_dtype not in ("fp32", "fp16", "bf16") or self.adapter_dtype not in ("fp32", "fp16", "bf16") or self.optimizer not in ("adamw-fp32", "sgd-no-momentum"): raise ValueError("unsupported dtype/optimizer; quantization overhead is not modeled") if type(self.held_out) is not bool or not self.held_out or self.eval_cases < self.min_eval_cases: raise ValueError("qualified held-out evidence required") score(self.quality) score(self.required_quality) for name in ("scope", "model", "data"): if getattr(self, "eval_" + name) != getattr(self, name): raise ValueError("evaluation binding mismatch") for name in ("adapter_ids", "target_ids"): value = getattr(self, name) if type(value) not in (list, tuple) or not 1 <= len(value) <= 10000: raise ValueError("bounded identity sequence required") for item in value: text(item) if len(set(value)) != len(value): raise ValueError("duplicate identity") object.__setattr__(self, name, tuple(value)) if type(self.targets) not in (tuple, list) or not 1 <= len(self.targets) <= 10000: raise ValueError("bounded targets required") copied = [] for item in self.targets: if type(item) is not Target: raise ValueError("concrete Target required") copied.append(Target(item.target_id, item.input_dim, item.output_dim, item.rank)) object.__setattr__(self, "targets", tuple(copied)) if tuple(item.target_id for item in copied) != self.target_ids: raise ValueError("target identity/order mismatch") if self.base_parameters % self.shard_count or any(item.input_dim % self.shard_count or item.output_dim % self.shard_count for item in copied): raise ValueError("teaching sharding requires exact divisibility") if sum(item.input_dim * item.output_dim for item in copied) > self.base_parameters: raise ValueError("targets exceed bound base model") if self.eval_config_digest != digest(configuration(self)): raise ValueError("evaluation configuration changed") def configuration(contract): return {name: ([asdict(x) for x in contract.targets] if name == "targets" else getattr(contract, name)) for name in ("policy", "version", "base_parameters", "base_dtype", "adapter_dtype", "optimizer", "adapter_ids", "target_ids", "shard_count", "targets")} @dataclass(frozen=True) class Budget: contract_digest: str trainable_parameters: int base_weight_bytes: int adapter_state_bytes: int lower_bound_bytes: int decision: str claim: str def audit(contract): if type(contract) is not BudgetContract: raise ValueError("concrete BudgetContract required") contract = BudgetContract(**{f.name: getattr(contract, f.name) for f in fields(BudgetContract)}) dtype_bytes = {"fp32": 4, "fp16": 2, "bf16": 2} trainable = sum(item.rank * (item.input_dim + item.output_dim) for item in contract.targets) * len(contract.adapter_ids) base_bytes = contract.base_parameters * dtype_bytes[contract.base_dtype] # Adapter weights + same-dtype gradients + optional FP32 master weights and Adam moments. state_per_parameter = 2 * dtype_bytes[contract.adapter_dtype] if contract.optimizer == "adamw-fp32": state_per_parameter += 8 + (4 if contract.adapter_dtype != "fp32" else 0) adapter_bytes = trainable * state_per_parameter return Budget(digest(asdict(contract)), trainable, base_bytes, adapter_bytes, base_bytes + adapter_bytes, "LOW_RANK_SUFFICIENT_ON_BOUND_EVAL" if contract.quality >= contract.required_quality else "BLOCK_LOW_RANK_INSUFFICIENT", "LOWER_BOUND_EXCLUDES_ACTIVATIONS_WORKSPACES_RUNTIME") def example_contract(): values = dict(scope="illustrative-support", model="base-v1", data="fixture-v1", policy="budget-v1", version="1", provenance="synthetic fixture", eval_digest=digest({"held_out_quality": 0.8}), eval_scope="illustrative-support", eval_model="base-v1", eval_data="fixture-v1", eval_cases=100, min_eval_cases=100, held_out=True, quality=0.8, required_quality=0.9, base_parameters=1_000_000, base_dtype="bf16", adapter_dtype="bf16", optimizer="adamw-fp32", adapter_ids=("support-v1",), target_ids=("layer.0.q_proj", "layer.0.v_proj"), shard_count=2, targets=(Target("layer.0.q_proj", 128, 128, 8), Target("layer.0.v_proj", 128, 128, 8))) config = {name: ([asdict(x) for x in values[name]] if name == "targets" else values[name]) for name in ("policy", "version", "base_parameters", "base_dtype", "adapter_dtype", "optimizer", "adapter_ids", "target_ids", "shard_count", "targets")} return BudgetContract(**values, eval_config_digest=digest(config)) def main(): result = audit(example_contract()) print("example=illustrative_only") print(f"trainable_parameters={result.trainable_parameters}") print(f"lower_bound_bytes={result.lower_bound_bytes}") print("decision=" + result.decision) print("claim=" + result.claim) if __name__ == "__main__": main()