"""Deterministic accelerator memory-traffic bound for an invented device.""" from dataclasses import asdict, dataclass, fields import hashlib import json import re IDENTITY = re.compile(r"[A-Za-z0-9][A-Za-z0-9_.:/@-]{0,127}") SHA256 = re.compile(r"[a-f0-9]{64}") TIER_NAMES = ("sram", "hbm", "interconnect", "host") def identity(value): if type(value) is not str or not IDENTITY.fullmatch(value): raise ValueError("exact bounded identity required") return value def integer(value, lower, upper): if type(value) is not int or not lower <= value <= upper: raise ValueError("bounded non-boolean integer required") return value def sequence(value, lower, upper): if type(value) not in (tuple, list): raise ValueError("bounded tuple or list required") integer(len(value), lower, upper) return tuple(value) def identities(value, lower, upper): result = sequence(value, lower, upper) for item in result: identity(item) if len(set(result)) != len(result): raise ValueError("duplicate identity") return result def scope_tuple(value): return identities(value, 8, 8) def sha256_text(value): if type(value) is not str or not SHA256.fullmatch(value): raise ValueError("exact lowercase SHA-256 required") return value def digest(value): encoded = json.dumps( value, sort_keys=True, separators=(",", ":"), allow_nan=False ).encode() return hashlib.sha256(encoded).hexdigest() def seal(record): if type(record.content_id) is not str: raise ValueError("content digest requires exact string") data = asdict(record) data.pop("content_id") return digest(data) def validate_record(record, cls): if type(record) is not cls: raise ValueError("concrete frozen record required") try: rebuilt = cls(**{field.name: getattr(record, field.name) for field in fields(cls)}) except (AttributeError, TypeError) as error: raise ValueError("malformed record") from error if record != rebuilt or record.content_id != rebuilt.content_id: raise ValueError("noncanonical or modified record") return rebuilt def ceil_div(numerator, denominator): integer(numerator, 0, 10**30) integer(denominator, 1, 10**30) return (numerator + denominator - 1) // denominator @dataclass(frozen=True) class MemoryTier: scope: tuple name: str capacity_bytes: int bandwidth_bytes_per_second: int content_id: str = "" def __post_init__(self): object.__setattr__(self, "scope", scope_tuple(self.scope)) if type(self.name) is not str or self.name not in TIER_NAMES: raise ValueError("unknown exact memory tier") integer(self.capacity_bytes, 1, 10**18) integer(self.bandwidth_bytes_per_second, 1, 10**18) expected = seal(self) if self.content_id and self.content_id != expected: raise ValueError("memory tier digest mismatch") object.__setattr__(self, "content_id", expected) @dataclass(frozen=True) class MemoryTrafficContract: scope: tuple tiers: tuple hbm_resident_limit_permille: int target_tokens_per_second: int max_token_floor_us: int content_id: str = "" def __post_init__(self): object.__setattr__(self, "scope", scope_tuple(self.scope)) tiers = tuple( validate_record(item, MemoryTier) for item in sequence(self.tiers, len(TIER_NAMES), len(TIER_NAMES)) ) if any(tier.scope != self.scope for tier in tiers): raise ValueError("memory tier outside contract scope") if tuple(tier.name for tier in tiers) != TIER_NAMES: raise ValueError("memory tiers must be unique and canonically ordered") object.__setattr__(self, "tiers", tiers) integer(self.hbm_resident_limit_permille, 1, 1_000) integer(self.target_tokens_per_second, 1, 1_000_000) integer(self.max_token_floor_us, 1, 86_400_000_000) expected = seal(self) if self.content_id and self.content_id != expected: raise ValueError("memory traffic contract digest mismatch") object.__setattr__(self, "content_id", expected) @dataclass(frozen=True) class TrafficLeg: scope: tuple leg_id: str tier: str bytes_per_token: int content_id: str = "" def __post_init__(self): object.__setattr__(self, "scope", scope_tuple(self.scope)) identity(self.leg_id) if type(self.tier) is not str or self.tier not in TIER_NAMES: raise ValueError("unknown exact traffic tier") integer(self.bytes_per_token, 1, 10**18) expected = seal(self) if self.content_id and self.content_id != expected: raise ValueError("traffic leg digest mismatch") object.__setattr__(self, "content_id", expected) @dataclass(frozen=True) class MemoryWorkload: scope: tuple contract_content_id: str workload_id: str resident_hbm_bytes: int compute_floor_us: int traffic_legs: tuple content_id: str = "" def __post_init__(self): object.__setattr__(self, "scope", scope_tuple(self.scope)) sha256_text(self.contract_content_id) identity(self.workload_id) integer(self.resident_hbm_bytes, 1, 10**18) integer(self.compute_floor_us, 1, 86_400_000_000) legs = tuple( validate_record(item, TrafficLeg) for item in sequence(self.traffic_legs, 1, 256) ) if any(leg.scope != self.scope for leg in legs): raise ValueError("traffic leg outside workload scope") if len({leg.leg_id for leg in legs}) != len(legs): raise ValueError("duplicate traffic leg identity") object.__setattr__(self, "traffic_legs", legs) expected = seal(self) if self.content_id and self.content_id != expected: raise ValueError("memory workload digest mismatch") object.__setattr__(self, "content_id", expected) @dataclass(frozen=True) class TierTraffic: tier: str bytes_per_token: int optimistic_transfer_floor_us: int @dataclass(frozen=True) class MemoryTrafficReport: decision: str violations: tuple limiting_resource: str optimistic_token_floor_us: int optimistic_tokens_per_second_ceiling: int hbm_resident_limit_bytes: int tier_traffic: tuple evidence_id: str claim: str = "LOCAL_OPTIMISTIC_TRAFFIC_BOUND_NOT_DEVICE_BENCHMARK" def audit_memory_traffic( contract: MemoryTrafficContract, workload: MemoryWorkload ) -> MemoryTrafficReport: """Calculate an optimistic full-overlap bound without running a device.""" contract = validate_record(contract, MemoryTrafficContract) workload = validate_record(workload, MemoryWorkload) if workload.scope != contract.scope or workload.contract_content_id != contract.content_id: raise ValueError("workload belongs to another memory traffic contract") tier_by_name = {tier.name: tier for tier in contract.tiers} bytes_by_tier = {name: 0 for name in TIER_NAMES} for leg in workload.traffic_legs: bytes_by_tier[leg.tier] += leg.bytes_per_token traffic = [] resources = [("compute", workload.compute_floor_us)] for name in TIER_NAMES: transfer_us = ceil_div( bytes_by_tier[name] * 1_000_000, tier_by_name[name].bandwidth_bytes_per_second, ) traffic.append(TierTraffic(name, bytes_by_tier[name], transfer_us)) resources.append((name, transfer_us)) limiting_resource, optimistic_floor = max(resources, key=lambda item: item[1]) tokens_per_second_ceiling = 1_000_000 // optimistic_floor hbm_resident_limit = ( tier_by_name["hbm"].capacity_bytes * contract.hbm_resident_limit_permille ) // 1_000 violations = [] if workload.resident_hbm_bytes > hbm_resident_limit: violations.append("hbm-residency-headroom") if optimistic_floor > contract.max_token_floor_us: violations.append("token-floor-latency") if tokens_per_second_ceiling < contract.target_tokens_per_second: violations.append("target-throughput") decision = "REJECT_MODEL" if violations else "WITHIN_DECLARED_BOUND" evidence_id = digest( { "contract": contract.content_id, "workload": workload.content_id, "decision": decision, "violations": violations, "limiting_resource": limiting_resource, "optimistic_floor_us": optimistic_floor, } ) return MemoryTrafficReport( decision, tuple(violations), limiting_resource, optimistic_floor, tokens_per_second_ceiling, hbm_resident_limit, tuple(traffic), evidence_id, ) def illustrative_fixture(): scope = ( "traffic-model-v1", "device-profile-v1", "kernel-profile-v1", "weights-v1", "kv-layout-v1", "interconnect-v1", "host-path-v1", "fixture-v1", ) tiers = ( MemoryTier(scope, "sram", 64_000_000, 20_000_000_000_000), MemoryTier(scope, "hbm", 80_000_000_000, 2_000_000_000_000), MemoryTier(scope, "interconnect", 32_000_000_000, 400_000_000_000), MemoryTier(scope, "host", 512_000_000_000, 64_000_000_000), ) contract = MemoryTrafficContract(scope, tiers, 800, 100, 10_000) legs = ( TrafficLeg(scope, "attention-tile", "sram", 64_000_000), TrafficLeg(scope, "weights-and-kv", "hbm", 16_000_000_000), TrafficLeg(scope, "tensor-exchange", "interconnect", 400_000_000), TrafficLeg(scope, "host-control-path", "host", 64_000_000), ) workload = MemoryWorkload( scope, contract.content_id, "decode-profile-invented-v1", 48_000_000_000, 6_000, legs, ) return contract, workload def main(): report = audit_memory_traffic(*illustrative_fixture()) print("example=illustrative_only") print(f"decision={report.decision}") print(f"limiting_resource={report.limiting_resource}") print(f"optimistic_token_floor_us={report.optimistic_token_floor_us}") print( "optimistic_tokens_per_second_ceiling=" f"{report.optimistic_tokens_per_second_ceiling}" ) print(f"claim={report.claim}") if __name__ == "__main__": main()