Skip to content

Commit 9c36597

Browse files
committed
fix(weights): handle empty colocated tensor buckets
1 parent 39a1580 commit 9c36597

2 files changed

Lines changed: 183 additions & 4 deletions

File tree

miles/backends/megatron_utils/update_weight/update_weight_from_tensor.py

Lines changed: 24 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -304,7 +304,7 @@ def _send_to_colocated_engine(
304304
long_live_tensors = []
305305

306306
if getattr(FlattenedTensorBucket, "supports_multi_dtypes", False):
307-
converted_named_tensors_by_dtypes = {"dtype": hf_named_tensors}
307+
converted_named_tensors_by_dtypes = {"dtype": hf_named_tensors} if hf_named_tensors else {}
308308
else:
309309
converted_named_tensors_by_dtypes = {}
310310
for name, tensor in hf_named_tensors:
@@ -354,13 +354,33 @@ def _send_to_colocated_engine(
354354
)
355355

356356
else:
357-
num_dtypes = len(serialized_named_tensors[0])
358-
for i in range(num_dtypes):
357+
num_buckets = max(len(tensors) for tensors in serialized_named_tensors)
358+
empty_serialized_tensor = None
359+
for i in range(num_buckets):
360+
serialized_tensors_for_bucket = []
361+
for tensors in serialized_named_tensors:
362+
if i < len(tensors):
363+
serialized_tensors_for_bucket.append(tensors[i])
364+
continue
365+
366+
if empty_serialized_tensor is None:
367+
empty_tensor_data = _empty_flattened_tensor_data()
368+
long_live_tensors.append(empty_tensor_data)
369+
empty_serialized_tensor = MultiprocessingSerializer.serialize(empty_tensor_data, output_str=True)
370+
serialized_tensors_for_bucket.append(empty_serialized_tensor)
371+
359372
kwargs = {
360-
"serialized_named_tensors": [tensors[i] for tensors in serialized_named_tensors],
373+
"serialized_named_tensors": serialized_tensors_for_bucket,
361374
"load_format": "flattened_bucket",
362375
"weight_version": str(weight_version),
363376
}
364377
refs.append(ipc_engine.update_weights_from_tensor.remote(**kwargs))
365378

366379
return refs, long_live_tensors
380+
381+
382+
def _empty_flattened_tensor_data():
383+
return {
384+
"flattened_tensor": torch.empty(0, dtype=torch.uint8, device=torch.cuda.current_device()),
385+
"metadata": [],
386+
}
Lines changed: 159 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,159 @@
1+
import sys
2+
import types
3+
from types import SimpleNamespace
4+
5+
6+
class _FakeFlattenedTensorBucket:
7+
supports_multi_dtypes = True
8+
9+
def __init__(self, *, named_tensors=None, flattened_tensor=None, metadata=None):
10+
if named_tensors is not None:
11+
if not named_tensors:
12+
raise ValueError("Cannot create empty tensor bucket")
13+
self._flattened_tensor = ("flattened", tuple(name for name, _ in named_tensors))
14+
self._metadata = tuple(name for name, _ in named_tensors)
15+
return
16+
17+
self._flattened_tensor = flattened_tensor
18+
self._metadata = metadata
19+
20+
def get_flattened_tensor(self):
21+
return self._flattened_tensor
22+
23+
def get_metadata(self):
24+
return self._metadata
25+
26+
27+
class _FakeMultiprocessingSerializer:
28+
@staticmethod
29+
def serialize(value, output_str):
30+
assert output_str is True
31+
return value
32+
33+
34+
_fake_sglang = types.ModuleType("miles.backends.megatron_utils.sglang")
35+
_fake_sglang.FlattenedTensorBucket = _FakeFlattenedTensorBucket
36+
_fake_sglang.MultiprocessingSerializer = _FakeMultiprocessingSerializer
37+
sys.modules.setdefault("miles.backends.megatron_utils.sglang", _fake_sglang)
38+
39+
_fake_common = types.ModuleType("miles.backends.megatron_utils.update_weight.common")
40+
_fake_common._check_weight_sync_results = lambda *args, **kwargs: None
41+
_fake_common.begin_weight_update = lambda *args, **kwargs: None
42+
_fake_common.end_weight_update = lambda *args, **kwargs: None
43+
sys.modules.setdefault("miles.backends.megatron_utils.update_weight.common", _fake_common)
44+
45+
46+
class _FakeHfWeightIteratorBase:
47+
@staticmethod
48+
def create(*args, **kwargs):
49+
return None
50+
51+
52+
_fake_hf_weight_iterator_base = types.ModuleType(
53+
"miles.backends.megatron_utils.update_weight.hf_weight_iterator_base"
54+
)
55+
_fake_hf_weight_iterator_base.HfWeightIteratorBase = _FakeHfWeightIteratorBase
56+
sys.modules.setdefault(
57+
"miles.backends.megatron_utils.update_weight.hf_weight_iterator_base",
58+
_fake_hf_weight_iterator_base,
59+
)
60+
61+
_fake_broadcast = types.ModuleType(
62+
"miles.backends.megatron_utils.update_weight.update_weight_from_distributed.broadcast"
63+
)
64+
_fake_broadcast.connect_rollout_engines_from_distributed = lambda *args, **kwargs: None
65+
_fake_broadcast.disconnect_rollout_engines_from_distributed = lambda *args, **kwargs: None
66+
_fake_broadcast.update_weights_from_distributed = lambda *args, **kwargs: []
67+
sys.modules.setdefault(
68+
"miles.backends.megatron_utils.update_weight.update_weight_from_distributed.broadcast",
69+
_fake_broadcast,
70+
)
71+
72+
from miles.backends.megatron_utils.update_weight import update_weight_from_tensor as update_weight # noqa: E402
73+
74+
75+
class _FakeRemoteMethod:
76+
def __init__(self):
77+
self.calls = []
78+
79+
def remote(self, **kwargs):
80+
self.calls.append(kwargs)
81+
return f"ref-{len(self.calls)}"
82+
83+
84+
class _FakeEngine:
85+
def __init__(self):
86+
self.update_weights_from_tensor = _FakeRemoteMethod()
87+
88+
89+
def _install_fakes(monkeypatch, gathered):
90+
state = SimpleNamespace(local_object=None)
91+
92+
def gather_object(obj, object_gather_list, dst, group):
93+
state.local_object = obj
94+
if object_gather_list is not None:
95+
object_gather_list[:] = gathered
96+
97+
fake_dist = SimpleNamespace(
98+
get_rank=lambda: 0,
99+
get_world_size=lambda group=None: len(gathered),
100+
gather_object=gather_object,
101+
)
102+
monkeypatch.setattr(update_weight, "dist", fake_dist)
103+
monkeypatch.setattr(update_weight, "FlattenedTensorBucket", _FakeFlattenedTensorBucket)
104+
monkeypatch.setattr(update_weight, "MultiprocessingSerializer", _FakeMultiprocessingSerializer)
105+
monkeypatch.setattr(update_weight.torch.cuda, "current_device", lambda: "cuda:0")
106+
monkeypatch.setattr(
107+
update_weight.torch,
108+
"empty",
109+
lambda size, dtype, device: {"size": size, "dtype": dtype, "device": device},
110+
)
111+
112+
return state
113+
114+
115+
def test_empty_colocated_bucket_still_participates_in_gather(monkeypatch):
116+
state = _install_fakes(monkeypatch, gathered=[[], []])
117+
engine = _FakeEngine()
118+
119+
refs, long_lived_tensors = update_weight._send_to_colocated_engine(
120+
[],
121+
ipc_engine=engine,
122+
ipc_gather_src=0,
123+
ipc_gather_group=object(),
124+
weight_version=3,
125+
)
126+
127+
assert state.local_object == []
128+
assert refs == []
129+
assert long_lived_tensors == []
130+
assert engine.update_weights_from_tensor.calls == []
131+
132+
133+
def test_source_rank_pads_empty_colocated_bucket_entries(monkeypatch):
134+
remote_serialized_bucket = {"flattened_tensor": ("remote",), "metadata": ("remote_weight",)}
135+
state = _install_fakes(monkeypatch, gathered=[[], [remote_serialized_bucket]])
136+
engine = _FakeEngine()
137+
138+
refs, long_lived_tensors = update_weight._send_to_colocated_engine(
139+
[],
140+
ipc_engine=engine,
141+
ipc_gather_src=0,
142+
ipc_gather_group=object(),
143+
weight_version=7,
144+
)
145+
146+
assert state.local_object == []
147+
assert refs == ["ref-1"]
148+
assert len(long_lived_tensors) == 1
149+
empty_bucket = long_lived_tensors[0]
150+
assert empty_bucket["metadata"] == []
151+
assert empty_bucket["flattened_tensor"] == {"size": 0, "dtype": update_weight.torch.uint8, "device": "cuda:0"}
152+
153+
assert engine.update_weights_from_tensor.calls == [
154+
{
155+
"serialized_named_tensors": [empty_bucket, remote_serialized_bucket],
156+
"load_format": "flattened_bucket",
157+
"weight_version": "7",
158+
}
159+
]

0 commit comments

Comments
 (0)