Coverage for src/lilbee/providers/fleet/placement_spec.py: 100%
87 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-08-14 11:46 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-08-14 11:46 +0000
1"""User-authored manual placement spec for the multi-GPU fleet."""
3from __future__ import annotations
5import json
6from collections.abc import Mapping
7from dataclasses import dataclass, field
9from lilbee.providers.roles import WorkerRole
11_KEY_DEVICES = "devices"
12_KEY_TENSOR_SPLIT = "tensor_split"
13_KEY_REPLICAS = "replicas"
16class PlacementError(ValueError):
17 """A placement spec is malformed or does not fit the hardware."""
20@dataclass(frozen=True)
21class RolePlacement:
22 """One role's manual placement: device pins, optional split, replica count."""
24 devices: tuple[int, ...]
25 tensor_split: tuple[int, ...] | None = None
26 replicas: int = 1
29@dataclass(frozen=True)
30class PlacementSpec:
31 """A manual placement for every active role, keyed by role."""
33 roles: Mapping[WorkerRole, RolePlacement] = field(default_factory=dict)
35 def to_json(self) -> str:
36 """Serialize to a compact JSON string keyed by role value."""
37 out: dict[str, dict[str, object]] = {}
38 for role, rp in self.roles.items():
39 entry: dict[str, object] = {_KEY_DEVICES: list(rp.devices)}
40 if rp.tensor_split is not None:
41 entry[_KEY_TENSOR_SPLIT] = list(rp.tensor_split)
42 if rp.replicas != 1:
43 entry[_KEY_REPLICAS] = rp.replicas
44 out[role.value] = entry
45 return json.dumps(out, sort_keys=True)
47 def __str__(self) -> str:
48 return self.to_json()
50 @classmethod
51 def from_json(cls, raw: str) -> PlacementSpec:
52 """Parse a JSON string produced by ``to_json`` into a ``PlacementSpec``."""
53 try:
54 data = json.loads(raw)
55 except (ValueError, TypeError) as exc:
56 raise PlacementError(f"placement is not valid JSON: {exc}") from exc
57 if not isinstance(data, dict):
58 raise PlacementError("placement must be a JSON object keyed by role")
59 roles: dict[WorkerRole, RolePlacement] = {}
60 for key, entry in data.items():
61 role = _role_for(key)
62 roles[role] = _role_placement(role, entry)
63 return cls(roles=roles)
66_ALLOWED_KEYS = frozenset({_KEY_DEVICES, _KEY_TENSOR_SPLIT, _KEY_REPLICAS})
69def _role_for(key: str) -> WorkerRole:
70 try:
71 return WorkerRole(key)
72 except ValueError as exc:
73 raise PlacementError(f"unknown role {key!r} in placement") from exc
76def _coerce_int(role: WorkerRole, field: str, value: object) -> int:
77 if isinstance(value, bool) or not isinstance(value, (int, float, str)):
78 raise PlacementError(f"{role.value}: {field} must be integers, got {value!r}")
79 if isinstance(value, float) and not value.is_integer():
80 raise PlacementError(f"{role.value}: {field} must be integers, got {value!r}")
81 try:
82 return int(value)
83 except (TypeError, ValueError) as exc:
84 raise PlacementError(f"{role.value}: {field} must be integers, got {value!r}") from exc
87def _int_list(role: WorkerRole, field: str, value: object) -> tuple[int, ...]:
88 if not isinstance(value, (list, tuple)):
89 raise PlacementError(f"{role.value}: {field} must be a list, got {type(value).__name__}")
90 return tuple(_coerce_int(role, field, item) for item in value)
93def _role_placement(role: WorkerRole, entry: object) -> RolePlacement:
94 if not isinstance(entry, dict):
95 raise PlacementError(f"{role.value}: placement entry must be an object")
96 unknown = set(entry) - _ALLOWED_KEYS
97 if unknown:
98 allowed = ", ".join(sorted(_ALLOWED_KEYS))
99 raise PlacementError(
100 f"{role.value}: unknown placement key(s) {sorted(unknown)}; allowed: {allowed}"
101 )
102 devices = _int_list(role, _KEY_DEVICES, entry.get(_KEY_DEVICES, []))
103 if not devices:
104 raise PlacementError(f"{role.value}: at least one device is required")
105 if any(d < 0 for d in devices):
106 raise PlacementError(f"{role.value}: device indices must be >= 0, got {list(devices)}")
107 if len(set(devices)) != len(devices):
108 raise PlacementError(f"{role.value}: duplicate device indices in {list(devices)}")
109 raw_split = entry.get(_KEY_TENSOR_SPLIT)
110 tensor_split: tuple[int, ...] | None = None
111 if raw_split is not None:
112 tensor_split = _int_list(role, _KEY_TENSOR_SPLIT, raw_split)
113 if len(tensor_split) != len(devices):
114 raise PlacementError(
115 f"{role.value}: tensor_split has {len(tensor_split)} weights "
116 f"for {len(devices)} devices"
117 )
118 if any(w <= 0 for w in tensor_split):
119 raise PlacementError(
120 f"{role.value}: tensor_split weights must be > 0, got {list(tensor_split)}"
121 )
122 replicas = _coerce_int(role, _KEY_REPLICAS, entry.get(_KEY_REPLICAS, 1))
123 if replicas < 1:
124 raise PlacementError(f"{role.value}: replicas must be >= 1")
125 return RolePlacement(devices=devices, tensor_split=tensor_split, replicas=replicas)