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

1"""User-authored manual placement spec for the multi-GPU fleet.""" 

2 

3from __future__ import annotations 

4 

5import json 

6from collections.abc import Mapping 

7from dataclasses import dataclass, field 

8 

9from lilbee.providers.roles import WorkerRole 

10 

11_KEY_DEVICES = "devices" 

12_KEY_TENSOR_SPLIT = "tensor_split" 

13_KEY_REPLICAS = "replicas" 

14 

15 

16class PlacementError(ValueError): 

17 """A placement spec is malformed or does not fit the hardware.""" 

18 

19 

20@dataclass(frozen=True) 

21class RolePlacement: 

22 """One role's manual placement: device pins, optional split, replica count.""" 

23 

24 devices: tuple[int, ...] 

25 tensor_split: tuple[int, ...] | None = None 

26 replicas: int = 1 

27 

28 

29@dataclass(frozen=True) 

30class PlacementSpec: 

31 """A manual placement for every active role, keyed by role.""" 

32 

33 roles: Mapping[WorkerRole, RolePlacement] = field(default_factory=dict) 

34 

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) 

46 

47 def __str__(self) -> str: 

48 return self.to_json() 

49 

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) 

64 

65 

66_ALLOWED_KEYS = frozenset({_KEY_DEVICES, _KEY_TENSOR_SPLIT, _KEY_REPLICAS}) 

67 

68 

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 

74 

75 

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 

85 

86 

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) 

91 

92 

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)