Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 27 additions & 0 deletions tests/utils/test_update_weight_from_distributed.py
Original file line number Diff line number Diff line change
Expand Up @@ -283,6 +283,11 @@ def test_update_uses_native_main_and_draft_lifecycles(update_module, monkeypatch
assert len(engine.start_draft_weight_update.calls) == 1
assert len(engine.finish_weight_update.calls) == 2
assert len(engine.continue_generation.calls) == 1
metrics = updater.pop_metrics()
assert metrics["weight_update_total_seconds"] >= metrics["weight_update_transfer_phase_seconds"]
assert metrics["weight_update_prepare_seconds"] >= 0
assert metrics["weight_update_finish_seconds"] >= 0
assert updater.pop_metrics() == {}


@pytest.mark.unit
Expand Down Expand Up @@ -349,6 +354,28 @@ def test_vllm_weight_iterator_keeps_checkpoint_scale_layout(weight_modules):
assert inspect.signature(direct_module.HfWeightIteratorDirect).parameters["transform_ue8m0"].default is False


@pytest.mark.unit
def test_weight_source_counts_lazy_chunks_once_per_update(weight_modules):
common, _ = weight_modules
tensor = torch.zeros(4, dtype=torch.float16)

class LazyIterator:
def get_hf_weight_chunks(self, weights):
yield ((name, value) for name, value in weights.items())

source = common.HfWeightSource(LazyIterator(), lambda: {"weight": tensor})
assert list(source) == [("weight", tensor)]
metrics = source.stop_metrics()
assert metrics["weight_update_bytes"] == 8
assert metrics["weight_update_chunks"] == 1
assert list(source) == [("weight", tensor)]
assert source.stop_metrics() == metrics

source.reset_metrics()
assert list(source) == [("weight", tensor)]
assert source.stop_metrics()["weight_update_chunks"] == 1


def _param_info(name: str, param: torch.Tensor, src_rank: int = 0) -> ParamInfo:
return ParamInfo(name, param.dtype, param.shape, {}, param.nbytes, src_rank)

Expand Down
27 changes: 27 additions & 0 deletions tests/utils/test_update_weight_from_tensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -269,6 +269,11 @@ def test_native_update_runs_main_and_draft_lifecycles(update_module, monkeypatch
assert len(engine.start_draft_weight_update.calls) == 1
assert len(engine.finish_weight_update.calls) == 2
assert len(engine.continue_generation.calls) == 1
metrics = updater.pop_metrics()
assert metrics["weight_update_total_seconds"] >= metrics["weight_update_transfer_phase_seconds"]
assert metrics["weight_update_prepare_seconds"] >= 0
assert metrics["weight_update_finish_seconds"] >= 0
assert updater.pop_metrics() == {}


@pytest.mark.unit
Expand Down Expand Up @@ -308,6 +313,28 @@ def test_failed_native_update_does_not_resume_generation(update_module, monkeypa
assert engine.continue_generation.calls == []


@pytest.mark.unit
def test_manual_transfer_materializes_lazy_chunk_before_send(update_module, monkeypatch):
updater = _updater(update_module)
tensor = torch.zeros(4, dtype=torch.float16)
updater._perf_export_seconds = 0.0
updater._perf_transferred_bytes = 0
updater._perf_chunk_count = 0
updater._hf_weight_iterator.get_hf_weight_chunks.return_value = iter(
[((name, value) for name, value in {"weight": tensor}.items())]
)
sent = []
monkeypatch.setattr(updater, "_send_hf_params", lambda chunk: (sent.append(chunk) or [], None))
monkeypatch.setattr(update_module.accelerator, "ipc_collect", lambda: None)
monkeypatch.setattr(update_module.accelerator, "empty_cache", lambda: None)

updater._send_weight_chunks({"weight": tensor})

assert sent == [[("weight", tensor)]]
assert updater._perf_transferred_bytes == 8
assert updater._perf_chunk_count == 1


@pytest.mark.unit
def test_native_ipc_buffer_covers_largest_reconstructed_tensor(update_module, monkeypatch):
dense = ParamInfo("dense", torch.float16, (8,), {}, 16, 0)
Expand Down
48 changes: 44 additions & 4 deletions vime/backends/megatron_utils/update_weight/common.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import inspect
import re
import socket
import time
from argparse import Namespace
from collections.abc import Callable, Iterator, Mapping, Sequence
from typing import Any
Expand Down Expand Up @@ -224,21 +225,54 @@ def __init__(
self.draft_weights_getter = draft_weights_getter
self.draft = False
self._metadata = {}
self.reset_metrics()

def reset_metrics(self) -> None:
self.export_seconds = 0.0
self.transferred_bytes = 0
self.chunk_count = 0
self._record_metrics = True

def stop_metrics(self) -> dict[str, float]:
self._record_metrics = False
return {
"weight_update_export_seconds": self.export_seconds,
"weight_update_bytes": float(self.transferred_bytes),
"weight_update_chunks": float(self.chunk_count),
}

def metadata(self):
if self.draft not in self._metadata:
from vllm.distributed.weight_transfer.base import ParamMeta

self._metadata[self.draft] = [ParamMeta(name, tensor.dtype, tuple(tensor.shape)) for name, tensor in self]
recording = self._record_metrics
self._record_metrics = False
try:
self._metadata[self.draft] = [
ParamMeta(name, tensor.dtype, tuple(tensor.shape)) for name, tensor in self
]
finally:
self._record_metrics = recording
return self._metadata[self.draft]

def __iter__(self):
if self.draft:
if self.draft_weights_getter is None:
raise RuntimeError("Draft weight update requested without a draft weight source")
yield from self.draft_weights_getter()
return
for chunk in self.iterator.get_hf_weight_chunks(self.weights_getter()):
chunks = ([item] for item in self.draft_weights_getter())
else:
chunks = self.iterator.get_hf_weight_chunks(self.weights_getter())
iterator = iter(chunks)
while True:
started = time.perf_counter()
try:
chunk = list(next(iterator))
except StopIteration:
break
if self._record_metrics:
self.export_seconds += time.perf_counter() - started
self.transferred_bytes += sum(tensor.numel() * tensor.element_size() for _, tensor in chunk)
self.chunk_count += 1
yield from chunk


Expand All @@ -253,6 +287,8 @@ def __init__(
self.version_getter = version_getter
self.engine_gpu_counts = engine_gpu_counts
self.draft = False
self.prepare_seconds = 0.0
self.finish_seconds = 0.0

def init_weight_transfer_engine(self, init_info: dict[str, Any]) -> None:
import ray
Expand All @@ -270,8 +306,10 @@ def init_weight_transfer_engine(self, init_info: dict[str, Any]) -> None:
def start_weight_update(self) -> None:
import ray

started = time.perf_counter()
method = "start_draft_weight_update" if self.draft else "start_weight_update"
ray.get([getattr(engine, method).remote() for engine in self.engines])
self.prepare_seconds += time.perf_counter() - started

def update_weights(self, update_info: dict[str, Any] | list[dict[str, Any]]) -> None:
import ray
Expand All @@ -281,8 +319,10 @@ def update_weights(self, update_info: dict[str, Any] | list[dict[str, Any]]) ->
def finish_weight_update(self, weight_version: str | None = None) -> None:
import ray

started = time.perf_counter()
version = str(self.version_getter()) if weight_version is None else str(weight_version)
ray.get([engine.finish_weight_update.remote(weight_version=version) for engine in self.engines])
self.finish_seconds += time.perf_counter() - started


def create_nccl_trainer(
Expand Down
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import time
from argparse import Namespace
from collections.abc import Callable, Mapping, Sequence
from functools import partial
Expand Down Expand Up @@ -85,8 +86,16 @@ def pop_metrics(self) -> dict[str, float]:
@torch.no_grad()
def update_weights(self) -> None:
assert self._trainer is not None
self.update_weight_metrics = {}
if hasattr(self._source, "reset_metrics"):
self._source.reset_metrics()
client = self._trainer.client
client.prepare_seconds = 0.0
client.finish_seconds = 0.0
total_started = time.perf_counter()
self.weight_version += 1

pause_started = time.perf_counter()
if dist.get_rank() == 0:
ray.get([engine.pause_generation.remote() for engine in self.rollout_engines])
ray.get([engine.flush_cache.remote() for engine in self.rollout_engines])
Expand All @@ -97,8 +106,9 @@ def update_weights(self) -> None:
rollout_engines=self.rollout_engines,
)
dist.barrier(group=get_gloo_group())
pause_flush_seconds = time.perf_counter() - pause_started

client = self._trainer.client
transfer_started = time.perf_counter()
client.draft = False
self._trainer.send_weights()
update_draft = self.args.dspark_enabled or (
Expand All @@ -110,7 +120,9 @@ def update_weights(self) -> None:
self._trainer.send_weights()
self._source.draft = False
client.draft = False
transfer_phase_seconds = time.perf_counter() - transfer_started

resume_started = time.perf_counter()
if dist.get_rank() == 0:
if self.quantization_config and self.quantization_config["quant_method"] in ["compressed-tensors"]:
post_process_weights(
Expand All @@ -120,6 +132,32 @@ def update_weights(self) -> None:
)
ray.get([engine.continue_generation.remote() for engine in self.rollout_engines])
dist.barrier(group=get_gloo_group())
resume_seconds = time.perf_counter() - resume_started

source_metrics = self._source.stop_metrics() if hasattr(self._source, "stop_metrics") else {}
transferred_bytes = source_metrics.get("weight_update_bytes", 0.0)
transfer_load_seconds = max(
0.0,
transfer_phase_seconds
- source_metrics.get("weight_update_export_seconds", 0.0)
- client.prepare_seconds
- client.finish_seconds,
)
self.update_weight_metrics = {
"weight_update_total_seconds": time.perf_counter() - total_started,
"weight_update_pause_flush_seconds": pause_flush_seconds,
"weight_update_prepare_seconds": client.prepare_seconds,
"weight_update_transfer_phase_seconds": transfer_phase_seconds,
"weight_update_export_seconds": source_metrics.get("weight_update_export_seconds", 0.0),
"weight_update_transfer_load_seconds": transfer_load_seconds,
"weight_update_finish_seconds": client.finish_seconds,
"weight_update_resume_seconds": resume_seconds,
"weight_update_bytes": transferred_bytes,
"weight_update_chunks": source_metrics.get("weight_update_chunks", 0.0),
"weight_update_effective_gib_per_second": (
transferred_bytes / (1024**3) / transfer_phase_seconds if transfer_phase_seconds > 0 else 0.0
),
}


def post_process_weights(
Expand Down
Loading