From dc665c44d61c193fd83997bdda43a7dc8e86cd62 Mon Sep 17 00:00:00 2001 From: "chucai.dzq" Date: Tue, 22 Sep 2026 17:01:57 +0800 Subject: [PATCH 1/2] feat(transfer): add bounded colocate transport and portable IPC allocation --- awex/reader/nccl_reader.py | 15 +- awex/tests/test_cuda_ipc.py | 101 +++++ awex/tests/test_nccl_bounded_stream.py | 536 ++++++++++++++++++++++ awex/transfer/nccl_bounded_stream.py | 593 +++++++++++++++++++++++++ awex/util/cuda_ipc.py | 66 +++ 5 files changed, 1307 insertions(+), 4 deletions(-) create mode 100644 awex/tests/test_cuda_ipc.py create mode 100644 awex/tests/test_nccl_bounded_stream.py create mode 100644 awex/transfer/nccl_bounded_stream.py create mode 100644 awex/util/cuda_ipc.py diff --git a/awex/reader/nccl_reader.py b/awex/reader/nccl_reader.py index 76c5dba..ff12b35 100644 --- a/awex/reader/nccl_reader.py +++ b/awex/reader/nccl_reader.py @@ -269,14 +269,21 @@ def _init_reader_in_colocate_mode(self): self.training_params_meta, self.infer_to_train_device_mapping[self.transfer_rank], ) + self.colocate_transport = self.create_colocate_transport() + logger.info( + f"Initialized NCCL weights reader for rank {self.transfer_rank} in colocate mode" + ) + + def create_colocate_transport(self): + """Create the transport; subclasses may select a different implementation. + + Every reader in a colocate group must select the same implementation. + """ from awex.transfer.nccl_stream_batch import NcclColocateStreamBatchTransport - self.colocate_transport = NcclColocateStreamBatchTransport( + return NcclColocateStreamBatchTransport( self.transfer_rank, self.infer_world_size ) - logger.info( - f"Initialized NCCL weights reader for rank {self.transfer_rank} in colocate mode" - ) def pre_update_weights(self, step_id, **kwargs): pass diff --git a/awex/tests/test_cuda_ipc.py b/awex/tests/test_cuda_ipc.py new file mode 100644 index 0000000..1700643 --- /dev/null +++ b/awex/tests/test_cuda_ipc.py @@ -0,0 +1,101 @@ +# Licensed to the Awex developers under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + + +from contextlib import nullcontext + +import pytest +import torch + +from awex.util.cuda_ipc import cuda_ipc_allocation + + +@pytest.mark.parametrize("fail", [False, True]) +def test_ipc_allocation_restores_runtime_options_after_exit(monkeypatch, fail): + original = ( + "expandable_segments:True,garbage_collection_threshold:0.7," + "max_split_size_mb:128,roundup_power2_divisions:[256:1,512:2,>:4]" + ) + state = {"config": original} + monkeypatch.setenv("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:False") + monkeypatch.setattr( + torch.cuda.memory, + "_snapshot", + lambda: { + "allocator_settings": { + "PYTORCH_CUDA_ALLOC_CONF": state["config"], + "expandable_segments": "expandable_segments:True" in state["config"], + } + }, + ) + monkeypatch.setattr( + torch.cuda.memory, + "_set_allocator_settings", + lambda config: state.update(config=config), + ) + with pytest.raises(RuntimeError, match="packing failed") if fail else nullcontext(): + with cuda_ipc_allocation(): + assert "expandable_segments:False" in state["config"] + assert "expandable_segments:True" not in state["config"] + assert "garbage_collection_threshold:0.7" in state["config"] + assert "max_split_size_mb:128" in state["config"] + assert "roundup_power2_divisions:[256:1,512:2,>:4]" in state["config"] + with cuda_ipc_allocation(): + assert "expandable_segments:False" in state["config"] + if fail: + raise RuntimeError("packing failed") + assert state["config"] == original + + +def test_ipc_allocation_when_disabled_does_not_change_allocator(monkeypatch): + monkeypatch.setattr( + torch.cuda.memory, + "_snapshot", + lambda: {"allocator_settings": {"expandable_segments": False}}, + ) + + def unexpected(*args): + pytest.fail("Disabled allocator must not be changed") + + monkeypatch.setattr(torch.cuda.memory, "_set_allocator_settings", unexpected) + with cuda_ipc_allocation(): + pass + + +def test_ipc_allocation_restores_flag_omitted_from_last_runtime_update(monkeypatch): + state = {"expandable": True, "config": "garbage_collection_threshold:0.7"} + + def set_config(config): + state["config"] = config + if "expandable_segments:" in config: + state["expandable"] = "expandable_segments:True" in config + + monkeypatch.setattr( + torch.cuda.memory, + "_snapshot", + lambda: { + "allocator_settings": { + "PYTORCH_CUDA_ALLOC_CONF": state["config"], + "expandable_segments": state["expandable"], + } + }, + ) + monkeypatch.setattr(torch.cuda.memory, "_set_allocator_settings", set_config) + with cuda_ipc_allocation(): + assert not state["expandable"] + assert state["expandable"] + assert "garbage_collection_threshold:0.7" in state["config"] diff --git a/awex/tests/test_nccl_bounded_stream.py b/awex/tests/test_nccl_bounded_stream.py new file mode 100644 index 0000000..7c8046e --- /dev/null +++ b/awex/tests/test_nccl_bounded_stream.py @@ -0,0 +1,536 @@ +# Licensed to the Awex developers under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + + +from contextlib import nullcontext +from types import SimpleNamespace + +import pytest +import torch + +from awex.transfer import nccl_bounded_stream as colocate_transport +from awex.transfer.nccl_bounded_stream import ( + BoundedMemoryNcclColocateStreamBatchTransport, +) + + +@pytest.mark.parametrize("pending_slice_copy", [False, True]) +def test_bounded_transport_defers_send_clones_until_execution( + monkeypatch, pending_slice_copy +): + """Building an AWEX transfer plan retains views instead of model-sized clones.""" + from awex.transfer import nccl_stream_batch + from awex.util import device as device_util + + class _SourceTensor: + def __init__(self) -> None: + self.clone_calls = 0 + + def clone(self): + self.clone_calls += 1 + return self + + source = _SourceTensor() + source.ready = not pending_slice_copy + send_op = SimpleNamespace( + send_shard_meta=SimpleNamespace(name="weight"), + recv_rank=1, + ) + send_plan = SimpleNamespace(operations={1: [send_op]}) + recv_op = SimpleNamespace(recv_shard_meta=SimpleNamespace(name="weight")) + recv_plan = SimpleNamespace(operations={1: [recv_op]}) + recv_storage = torch.full((2, 4), float("nan")) + recv_target = recv_storage[:, ::2] + expected_recv = torch.arange(4, dtype=torch.float32).reshape(2, 2) + transport = object.__new__(BoundedMemoryNcclColocateStreamBatchTransport) + + def _inspect_plan( + transfer_rank, + world_size, + all_send_p2p_ops, + all_recv_p2p_ops, + weights_update_group, + rank_coordinate, + step_id, + ) -> None: + del ( + transfer_rank, + world_size, + weights_update_group, + rank_coordinate, + step_id, + ) + assert all_send_p2p_ops[1][0][1].tensor is source + assert source.clone_calls == 0 + # A materialized slice cannot be consumed on the transfer stream until + # its asynchronous producer on the caller stream has completed. + assert source.ready + recv_buffer = all_recv_p2p_ops[1][0][1].tensor + assert recv_buffer.is_contiguous() + assert recv_buffer.data_ptr() != recv_target.data_ptr() + recv_buffer.copy_(expected_recv) + + monkeypatch.setattr(transport, "_validate_pack_config", lambda group: None) + transport.execute_recursive_partition_stream_transfer = _inspect_plan + monkeypatch.setattr( + nccl_stream_batch, + "hang_detector", + SimpleNamespace(submit=lambda *args, **kwargs: None), + ) + monkeypatch.setattr( + "awex.transfer.nccl_comm.validate_rank_mappings", lambda *args: None + ) + monkeypatch.setattr( + "awex.transfer.transfer_plan.slice_tensor", + lambda tensor, *args, **kwargs: tensor, + ) + monkeypatch.setattr( + device_util, "synchronize", lambda: setattr(source, "ready", True) + ) + monkeypatch.setattr( + torch.distributed, + "P2POp", + lambda op, tensor, peer, group: SimpleNamespace( + op=op, tensor=tensor, peer=peer, group=group + ), + ) + + transport.update_weights_in_colocate_mode( + train_to_infer_device_mapping={0: 0, 1: 1}, + infer_to_train_device_mapping={0: 0, 1: 1}, + transfer_rank=0, + rank_coordinate="0-0-0", + world_size=2, + send_transfer_plan=send_plan, + recv_transfer_plan=recv_plan, + weights_update_group=object(), + send_parameters={"weight": source}, + recv_parameters={"weight": recv_target}, + step_id=1, + ) + + assert source.clone_calls == 0 + torch.testing.assert_close(recv_target, expected_recv, rtol=0, atol=0) + assert torch.isnan(recv_storage[:, 1::2]).all() + + +def test_bounded_transport_releases_each_send_clone_batch(monkeypatch): + """Only one send tensor per active peer remains live during P2P execution.""" + from awex.util import device as device_util + + counters = {"live": 0, "max_live": 0, "clones": 0, "syncs": 0} + + class _Clone: + def __init__(self) -> None: + counters["live"] += 1 + counters["max_live"] = max(counters["max_live"], counters["live"]) + + def __del__(self) -> None: + counters["live"] -= 1 + + class _SourceTensor: + def clone(self): + counters["clones"] += 1 + return _Clone() + + class _Work: + def __init__(self, tensor) -> None: + self.tensor = tensor + + def wait(self) -> None: + self.tensor = None + + def _isend(tensor, peer, group): + del peer, group + return _Work(tensor) + + monkeypatch.setattr(torch.distributed, "isend", _isend) + monkeypatch.setattr(device_util, "stream", lambda stream: nullcontext()) + monkeypatch.setattr( + device_util, + "synchronize", + lambda: counters.__setitem__("syncs", counters["syncs"] + 1), + ) + + transport = object.__new__(BoundedMemoryNcclColocateStreamBatchTransport) + transport._stream_pool = [object(), object()] + transport._expert_pack_config = (1, 64 * 1024 * 1024) + ops = { + peer: [ + ( + SimpleNamespace(recv_shard_meta=SimpleNamespace(dtype=None)), + SimpleNamespace( + op=_isend, + tensor=_SourceTensor(), + peer=peer, + group=object(), + ), + ) + for _ in range(3) + ] + for peer in (1, 2) + } + + count = transport._execute_ops_concurrent(ops, range(1, 3)) + + assert count == 6 + assert counters == { + "live": 0, + "max_live": 2, + "clones": 6, + "syncs": 3, + } + + +def test_bounded_transport_casts_send_to_receiver_dtype(monkeypatch): + """P2P sends use the receiver dtype so NCCL wire sizes match.""" + from awex.util import device as device_util + + sent = [] + + class _Work: + def wait(self) -> None: + return None + + def _isend(tensor, peer, group): + del peer, group + sent.append(tensor) + return _Work() + + source = torch.ones(2, dtype=torch.bfloat16) + plan_op = SimpleNamespace( + recv_shard_meta=SimpleNamespace(dtype=torch.float32), + ) + p2p_op = SimpleNamespace( + op=_isend, + tensor=source, + peer=1, + group=object(), + ) + + monkeypatch.setattr(torch.distributed, "isend", _isend) + monkeypatch.setattr(device_util, "stream", lambda stream: nullcontext()) + monkeypatch.setattr(device_util, "synchronize", lambda: None) + + transport = object.__new__(BoundedMemoryNcclColocateStreamBatchTransport) + transport._stream_pool = [object()] + transport._expert_pack_config = (1, 64 * 1024 * 1024) + + count = transport._execute_ops_concurrent( + {1: [(plan_op, p2p_op)]}, + range(1, 2), + ) + + assert count == 1 + assert len(sent) == 1 + assert sent[0].dtype == torch.float32 + + +def test_bounded_transport_prepares_send_on_transfer_stream(monkeypatch): + """Send clones are ordered on the same stream as their NCCL operation.""" + from awex.util import device as device_util + + state = {"active_stream": None} + transfer_stream = object() + + class _StreamContext: + def __enter__(self): + state["active_stream"] = transfer_stream + + def __exit__(self, exc_type, exc_value, traceback): + state["active_stream"] = None + + class _SourceTensor: + dtype = torch.bfloat16 + + def clone(self): + assert state["active_stream"] is transfer_stream + return self + + class _Work: + def wait(self) -> None: + return None + + def _isend(tensor, peer, group): + del tensor, peer, group + assert state["active_stream"] is transfer_stream + return _Work() + + monkeypatch.setattr(torch.distributed, "isend", _isend) + monkeypatch.setattr(device_util, "stream", lambda stream: _StreamContext()) + monkeypatch.setattr(device_util, "synchronize", lambda: None) + + transport = object.__new__(BoundedMemoryNcclColocateStreamBatchTransport) + transport._stream_pool = [transfer_stream] + transport._expert_pack_config = (1, 64 * 1024 * 1024) + plan_op = SimpleNamespace( + recv_shard_meta=SimpleNamespace(dtype=torch.bfloat16), + ) + p2p_op = SimpleNamespace( + op=_isend, + tensor=_SourceTensor(), + peer=1, + group=object(), + ) + + count = transport._execute_ops_concurrent( + {1: [(plan_op, p2p_op)]}, + range(1, 2), + ) + + assert count == 1 + assert state["active_stream"] is None + + +def test_expert_pack_defaults_to_validated_optimized_batch(monkeypatch): + monkeypatch.delenv("AWEX_EXPERT_PACK_OPS", raising=False) + monkeypatch.delenv("AWEX_EXPERT_PACK_MB", raising=False) + + assert colocate_transport._expert_pack_limits() == ( + 64, + 64 * 1024 * 1024, + ) + + +def test_bounded_transport_partitions_only_compatible_expert_ops(): + """Expert packs respect FIFO, byte limits, and non-expert boundaries.""" + group = object() + + def _send(tensor, peer, group): + del tensor, peer, group + + def _operation(param_class, dtype=torch.bfloat16): + plan_op = SimpleNamespace( + param_class=param_class, + overlap_shape=(2,), + recv_shard_meta=SimpleNamespace(dtype=dtype), + ) + p2p_op = SimpleNamespace( + op=_send, + tensor=torch.zeros(2, dtype=dtype), + peer=1, + group=group, + ) + return plan_op, p2p_op + + operations = [ + _operation("expert"), + _operation("expert"), + _operation("expert"), + _operation("dense_other"), + _operation("expert", torch.float32), + ] + + batches = ( + BoundedMemoryNcclColocateStreamBatchTransport._partition_expert_operations( + operations, + max_pack_ops=4, + max_pack_bytes=8, + ) + ) + + assert [len(batch) for batch in batches] == [2, 1, 1, 1] + flattened = [item for batch in batches for item in batch] + assert all( + actual_plan is expected_plan and actual_p2p is expected_p2p + for (actual_plan, actual_p2p), (expected_plan, expected_p2p) in zip( + flattened, operations + ) + ) + + +def test_bounded_transport_packs_expert_sends(monkeypatch): + """Consecutive expert sends share one flat wire tensor per bounded pack.""" + from awex.util import device as device_util + + sent = [] + syncs = [] + + class _Work: + def wait(self) -> None: + return None + + def _isend(tensor, peer, group): + del peer, group + sent.append(tensor.clone()) + return _Work() + + monkeypatch.setattr(torch.distributed, "isend", _isend) + monkeypatch.setattr(device_util, "stream", lambda stream: nullcontext()) + monkeypatch.setattr(device_util, "synchronize", lambda: syncs.append(None)) + + group = object() + source_tensors = [ + torch.tensor([1.0, 2.0]), + torch.tensor([[3.0, 4.0]]), + torch.tensor([5.0, 6.0]), + torch.tensor([[7.0, 8.0]]), + ] + operations = [] + for tensor in source_tensors: + plan_op = SimpleNamespace( + param_class="expert", + overlap_shape=tuple(tensor.shape), + recv_shard_meta=SimpleNamespace(dtype=torch.float32), + ) + p2p_op = SimpleNamespace( + op=_isend, + tensor=tensor, + peer=1, + group=group, + ) + operations.append((plan_op, p2p_op)) + + transport = object.__new__(BoundedMemoryNcclColocateStreamBatchTransport) + transport._stream_pool = [object()] + transport._expert_pack_config = (2, 64 * 1024 * 1024) + transport._expert_pack_stats = { + "logical_ops": 0, + "wire_ops": 0, + "packed_wire_ops": 0, + } + + count = transport._execute_ops_concurrent({1: operations}, range(1, 2)) + + assert count == 2 + assert len(syncs) == 2 + assert transport._expert_pack_stats == { + "logical_ops": 4, + "wire_ops": 2, + "packed_wire_ops": 2, + } + torch.testing.assert_close( + sent[0], torch.tensor([1.0, 2.0, 3.0, 4.0]), rtol=0.0, atol=0.0 + ) + torch.testing.assert_close( + sent[1], torch.tensor([5.0, 6.0, 7.0, 8.0]), rtol=0.0, atol=0.0 + ) + + +def test_bounded_transport_unpacks_expert_receives(monkeypatch): + """A flat expert receive is copied back into its original tensor views.""" + from awex.util import device as device_util + + payload = torch.tensor([1.0, 2.0, 3.0, 4.0]) + + class _Work: + def wait(self) -> None: + return None + + def _irecv(tensor, peer, group): + del peer, group + tensor.copy_(payload) + return _Work() + + monkeypatch.setattr(torch.distributed, "irecv", _irecv) + monkeypatch.setattr(device_util, "stream", lambda stream: nullcontext()) + monkeypatch.setattr(device_util, "synchronize", lambda: None) + + group = object() + destinations = [torch.zeros(2), torch.zeros(1, 2)] + operations = [] + for tensor in destinations: + plan_op = SimpleNamespace( + param_class="expert", + overlap_shape=tuple(tensor.shape), + recv_shard_meta=SimpleNamespace(dtype=torch.float32), + ) + p2p_op = SimpleNamespace( + op=_irecv, + tensor=tensor, + peer=1, + group=group, + ) + operations.append((plan_op, p2p_op)) + + transport = object.__new__(BoundedMemoryNcclColocateStreamBatchTransport) + transport._stream_pool = [object()] + transport._expert_pack_config = (2, 64 * 1024 * 1024) + transport._expert_pack_stats = None + + count = transport._execute_ops_concurrent({1: operations}, range(1, 2)) + + assert count == 1 + torch.testing.assert_close( + destinations[0], torch.tensor([1.0, 2.0]), rtol=0.0, atol=0.0 + ) + torch.testing.assert_close( + destinations[1], torch.tensor([[3.0, 4.0]]), rtol=0.0, atol=0.0 + ) + + +@pytest.mark.parametrize("mismatch", [False, True]) +def test_wire_config_rejects_rank_disagreement(monkeypatch, mismatch): + transport = object.__new__(BoundedMemoryNcclColocateStreamBatchTransport) + transport._expert_pack_config = (64, 64 * 1024 * 1024) + group = object() + calls = [] + monkeypatch.setattr(torch.distributed, "get_backend", lambda pg: "gloo") + + def reduce(tensor, op, group): + calls.append(group) + if mismatch and op == torch.distributed.ReduceOp.MIN: + tensor[0] = 1 + + monkeypatch.setattr(torch.distributed, "all_reduce", reduce) + expected = ( + pytest.raises(ValueError, match="differs across ranks") + if mismatch + else nullcontext() + ) + with expected: + transport._validate_pack_config(group) + assert calls == [group, group] + + +def _check_pack_config_worker(rank, rendezvous, mismatch): + from datetime import timedelta + + import torch.distributed as dist + + dist.init_process_group( + "gloo", + init_method=rendezvous, + rank=rank, + world_size=2, + timeout=timedelta(seconds=20), + ) + try: + transport = object.__new__(BoundedMemoryNcclColocateStreamBatchTransport) + transport._expert_pack_config = ( + 64 + (rank if mismatch else 0), + 64 * 1024 * 1024, + ) + if mismatch: + with pytest.raises(ValueError, match="differs across ranks"): + transport._validate_pack_config(dist.group.WORLD) + else: + transport._validate_pack_config(dist.group.WORLD) + finally: + dist.destroy_process_group() + + +@pytest.mark.parametrize("mismatch", [False, True]) +def test_pack_config_agrees_or_fails_on_all_real_gloo_ranks(tmp_path, mismatch): + import torch.multiprocessing as mp + + mp.spawn( + _check_pack_config_worker, + args=(f"file://{tmp_path / 'rendezvous'}", mismatch), + nprocs=2, + join=True, + ) diff --git a/awex/transfer/nccl_bounded_stream.py b/awex/transfer/nccl_bounded_stream.py new file mode 100644 index 0000000..76966bc --- /dev/null +++ b/awex/transfer/nccl_bounded_stream.py @@ -0,0 +1,593 @@ +# Licensed to the Awex developers under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + + +"""Bounded-memory colocate transfer with consistent expert packing across ranks.""" + +from __future__ import annotations + +import math +import os +import time +from typing import Any + +import torch +import torch.distributed as dist + +from awex.logging import getLogger +from awex.transfer.nccl_stream_batch import ( + NcclColocateStreamBatchTransport, +) + +logger = getLogger("BoundedColocateTransport") +_DEFAULT_EXPERT_PACK_OPS = 64 +_DEFAULT_EXPERT_PACK_MB = 64 + + +def _read_positive_env_int(name: str, default: int) -> int: + """Read a positive integer env value, preserving the safe default on errors.""" + raw_value = os.environ.get(name, "").strip() + if not raw_value: + return default + try: + return max(1, int(raw_value)) + except ValueError: + logger.warning("Ignoring invalid %s=%r; using %d", name, raw_value, default) + return default + + +def _expert_pack_limits() -> tuple[int, int]: + max_ops = _read_positive_env_int("AWEX_EXPERT_PACK_OPS", _DEFAULT_EXPERT_PACK_OPS) + max_bytes = ( + _read_positive_env_int("AWEX_EXPERT_PACK_MB", _DEFAULT_EXPERT_PACK_MB) + * 1024 + * 1024 + ) + return max_ops, max_bytes + + +class BoundedMemoryNcclColocateStreamBatchTransport(NcclColocateStreamBatchTransport): + """Run AWEX recursive P2P without retaining every send clone at once. + + Upstream AWEX clones every remote send slice while constructing the transfer + plan. A Qwen3-30B 8-way colocate update consequently retains roughly 7/8 + of the model (about 53 GiB per GPU) before NCCL starts. Keep source views + in the plan and materialize only one operation per active peer at a time. + The temporary clones stay alive until their sends complete, then become + reusable by the CUDA allocator before the next operation index. + + ``AWEX_EXPERT_PACK_OPS`` combines up to 64 consecutive routed-expert + operations for the same peer into one flat wire tensor by default. Set it + to one to restore the unpacked path. ``AWEX_EXPERT_PACK_MB`` bounds each + flat tensor and defaults to 64 MiB. + """ + + def _validate_pack_config(self, process_group) -> None: + """Reject inconsistent wire batching on every rank before issuing P2P.""" + backend = dist.get_backend(process_group) + device = ( + torch.device("cuda", torch.cuda.current_device()) + if backend == "nccl" + else torch.device("cpu") + ) + config = torch.tensor( + self._expert_pack_config, dtype=torch.int64, device=device + ) + minimum, maximum = config.clone(), config.clone() + dist.all_reduce(minimum, op=dist.ReduceOp.MIN, group=process_group) + dist.all_reduce(maximum, op=dist.ReduceOp.MAX, group=process_group) + if not torch.equal(minimum, maximum): + raise ValueError("AWEX expert packing configuration differs across ranks") + + def update_weights_in_colocate_mode( + self, + train_to_infer_device_mapping, + infer_to_train_device_mapping, + transfer_rank, + rank_coordinate, + world_size, + send_transfer_plan, + recv_transfer_plan, + weights_update_group, + send_parameters, + recv_parameters, + *, + step_id=-1, + async_op=True, + **kwargs, + ): + import os + from concurrent.futures import Future + + from awex.transfer.nccl_comm import ( + detect_hang, + execute_tensors_to_copy, + validate_rank_mappings, + ) + from awex.transfer.nccl_stream_batch import hang_detector + from awex.transfer.transfer_plan import slice_tensor + from awex.util import device as device_util + + logger.info( + "Using bounded-memory RECURSIVE PARTITION P2P for rank %s", + rank_coordinate, + ) + task_id = f"{rank_coordinate}-{step_id}" + validate_rank_mappings( + train_to_infer_device_mapping, infer_to_train_device_mapping + ) + start_time = time.time() + + self._expert_pack_config = _expert_pack_limits() + self._validate_pack_config(weights_update_group) + expert_pack_ops, expert_pack_bytes = self._expert_pack_config + self._expert_pack_stats = None + if expert_pack_ops > 1: + self._expert_pack_stats = { + "logical_ops": 0, + "wire_ops": 0, + "packed_wire_ops": 0, + } + logger.info( + "Expert P2P packing enabled for %s: max_ops=%d, max_mb=%.1f", + task_id, + expert_pack_ops, + expert_pack_bytes / 1024 / 1024, + ) + + send_ops = dict(send_transfer_plan.operations) + recv_ops = dict(recv_transfer_plan.operations) + num_sends = sum(len(ops) for ops in send_ops.values()) + num_recvs = sum(len(ops) for ops in recv_ops.values()) + logger.info( + "Start bounded-memory weights update for %s, num_sends=%d, num_recvs=%d", + task_id, + num_sends, + num_recvs, + ) + + all_send_p2p_ops = {} + all_recv_p2p_ops = {} + tensors_to_copy = [] + train_slice_context = {} + non_contiguous_tensor_pairs = [] + + for peer_rank, ops in send_ops.items(): + mapped_peer_rank = train_to_infer_device_mapping.get(peer_rank, peer_rank) + if mapped_peer_rank == transfer_rank: + for op in ops: + send_tensor = send_parameters[op.send_shard_meta.name] + tensor_sliced = slice_tensor( + send_tensor, + op, + True, + slice_context=train_slice_context, + ) + tensors_to_copy.append(tensor_sliced) + continue + + p2p_ops = [] + for op in ops: + send_tensor = send_parameters[op.send_shard_meta.name] + tensor_sliced = slice_tensor( + send_tensor, + op, + True, + slice_context=train_slice_context, + ) + recv_rank = train_to_infer_device_mapping.get( + op.recv_rank, op.recv_rank + ) + # Deliberately retain the source view. _execute_ops_concurrent + # clones a bounded batch immediately before enqueueing sends. + p2p_op = dist.P2POp( + dist.isend if async_op else dist.send, + tensor_sliced, + recv_rank, + group=weights_update_group, + ) + p2p_ops.append((op, p2p_op)) + all_send_p2p_ops[mapped_peer_rank] = p2p_ops + + for send_rank, ops in recv_ops.items(): + recv_from_rank = train_to_infer_device_mapping[send_rank] + if recv_from_rank == transfer_rank: + continue + p2p_ops = [] + for op in ops: + recv_tensor = recv_parameters[op.recv_shard_meta.name] + tensor_sliced = slice_tensor(recv_tensor, op, False) + if not tensor_sliced.is_contiguous(): + original_tensor = tensor_sliced + tensor_sliced = torch.empty_like( + tensor_sliced, memory_format=torch.contiguous_format + ) + non_contiguous_tensor_pairs.append((original_tensor, tensor_sliced)) + p2p_op = dist.P2POp( + dist.irecv if async_op else dist.recv, + tensor_sliced, + recv_from_rank, + group=weights_update_group, + ) + p2p_ops.append((op, p2p_op)) + all_recv_p2p_ops[recv_from_rank] = p2p_ops + + if tensors_to_copy: + send_rank = infer_to_train_device_mapping[transfer_rank] + execute_tensors_to_copy( + tensors_to_copy, + recv_transfer_plan.operations[send_rank], + recv_parameters, + f"tensor copy for {task_id}", + ) + else: + logger.info("No tensors to copy for %s", task_id) + + # slice_tensor may materialize send slices on the caller stream. + # Finish planning copies before independent transfer streams consume + # them, including ranks with no local copy to synchronize implicitly. + device_util.synchronize() + + future = Future() + total_send_ops = sum(len(ops) for ops in all_send_p2p_ops.values()) + total_recv_ops = sum(len(ops) for ops in all_recv_p2p_ops.values()) + message = ( + f"[{os.getpid()}] execute {total_send_ops} sends " + f"{total_recv_ops} recvs with bounded recursive partition for {task_id}" + ) + hang_detector.submit(detect_hang, future, message, [], timeout=60) + + self.execute_recursive_partition_stream_transfer( + transfer_rank, + world_size, + all_send_p2p_ops, + all_recv_p2p_ops, + weights_update_group, + rank_coordinate, + step_id, + ) + if non_contiguous_tensor_pairs: + with torch.no_grad(): + for original_tensor, recv_tensor in non_contiguous_tensor_pairs: + original_tensor.copy_(recv_tensor) + non_contiguous_tensor_pairs.clear() + device_util.synchronize() + future.set_result(True) + if self._expert_pack_stats is not None: + stats = self._expert_pack_stats + logger.info( + "Expert P2P packing finished for %s: logical_ops=%d, wire_ops=%d, " + "packed_wire_ops=%d", + task_id, + stats["logical_ops"], + stats["wire_ops"], + stats["packed_wire_ops"], + ) + logger.info( + "Finished bounded-memory weights update for %s, took %.4f seconds", + task_id, + time.time() - start_time, + ) + + @staticmethod + def _is_expert_operation(plan_op: Any) -> bool: + if getattr(plan_op, "param_class", None) == "expert": + return True + for meta_name in ("send_shard_meta", "recv_shard_meta"): + name = getattr(getattr(plan_op, meta_name, None), "name", "") + if ".experts." in name: + return True + return False + + @staticmethod + def _operation_wire_dtype(plan_op: Any, p2p_op: Any) -> torch.dtype: + recv_dtype = getattr(getattr(plan_op, "recv_shard_meta", None), "dtype", None) + return recv_dtype if recv_dtype is not None else p2p_op.tensor.dtype + + @staticmethod + def _operation_wire_numel(plan_op: Any, p2p_op: Any) -> int: + overlap_shape = getattr(plan_op, "overlap_shape", None) + if overlap_shape is not None: + return math.prod(overlap_shape) + return p2p_op.tensor.numel() + + @classmethod + def _partition_expert_operations( + cls, + operations: list[tuple[Any, Any]], + max_pack_ops: int, + max_pack_bytes: int, + ) -> list[list[tuple[Any, Any]]]: + """Pack consecutive compatible expert ops without changing FIFO order.""" + batches: list[list[tuple[Any, Any]]] = [] + current_batch: list[tuple[Any, Any]] = [] + current_signature = None + current_bytes = 0 + + def flush_current() -> None: + nonlocal current_batch, current_signature, current_bytes + if current_batch: + batches.append(current_batch) + current_batch = [] + current_signature = None + current_bytes = 0 + + for plan_op, p2p_op in operations: + if not cls._is_expert_operation(plan_op): + flush_current() + batches.append([(plan_op, p2p_op)]) + continue + + wire_dtype = cls._operation_wire_dtype(plan_op, p2p_op) + wire_bytes = ( + cls._operation_wire_numel(plan_op, p2p_op) * wire_dtype.itemsize + ) + signature = (p2p_op.op, p2p_op.peer, id(p2p_op.group), wire_dtype) + exceeds_limit = current_batch and ( + len(current_batch) >= max_pack_ops + or current_bytes + wire_bytes > max_pack_bytes + or signature != current_signature + ) + if exceeds_limit: + flush_current() + + if wire_bytes > max_pack_bytes: + batches.append([(plan_op, p2p_op)]) + continue + + current_batch.append((plan_op, p2p_op)) + current_signature = signature + current_bytes += wire_bytes + + flush_current() + return batches + + @classmethod + def _pack_send_batch(cls, batch: list[tuple[Any, Any]]) -> torch.Tensor: + wire_dtype = cls._operation_wire_dtype(*batch[0]) + source_tensors = [p2p_op.tensor for _, p2p_op in batch] + numels = [tensor.numel() for tensor in source_tensors] + packed = torch.empty( + sum(numels), + dtype=wire_dtype, + device=source_tensors[0].device, + ) + packed_views = [ + flat_view.view_as(source) + for flat_view, source in zip(packed.split(numels), source_tensors) + ] + with torch.no_grad(): + torch._foreach_copy_(packed_views, source_tensors) + return packed + + @staticmethod + def _allocate_packed_recv_batch( + batch: list[tuple[Any, Any]], wire_dtype: torch.dtype + ) -> torch.Tensor: + destination_tensors = [p2p_op.tensor for _, p2p_op in batch] + return torch.empty( + sum(tensor.numel() for tensor in destination_tensors), + dtype=wire_dtype, + device=destination_tensors[0].device, + ) + + @staticmethod + def _unpack_recv_batch(packed: torch.Tensor, batch: list[tuple[Any, Any]]) -> None: + destination_tensors = [p2p_op.tensor for _, p2p_op in batch] + numels = [tensor.numel() for tensor in destination_tensors] + packed_views = [ + flat_view.view_as(destination) + for flat_view, destination in zip(packed.split(numels), destination_tensors) + ] + with torch.no_grad(): + torch._foreach_copy_(destination_tensors, packed_views) + + def _execute_ops_concurrent(self, ops_dict, peer_ranks): + expert_pack_config = getattr(self, "_expert_pack_config", None) + if expert_pack_config is None: + expert_pack_config = _expert_pack_limits() + max_pack_ops, max_pack_bytes = expert_pack_config + if max_pack_ops <= 1: + return self._execute_ops_concurrent_unpacked(ops_dict, peer_ranks) + return self._execute_ops_concurrent_packed( + ops_dict, + peer_ranks, + max_pack_ops, + max_pack_bytes, + ) + + def _execute_ops_concurrent_unpacked(self, ops_dict, peer_ranks): + """Execute one tensor per active peer and release send clones promptly.""" + from awex.util import device as device_util + + peer_ops_with_rank = [ + (peer_rank, ops_dict[peer_rank]) + for peer_rank in peer_ranks + if peer_rank in ops_dict + ] + if not peer_ops_with_rank: + return 0 + + peer_to_stream_idx = { + peer_rank: index % len(self._stream_pool) + for index, (peer_rank, _) in enumerate(peer_ops_with_rank) + } + max_ops = max(len(ops) for _, ops in peer_ops_with_rank) + total_ops = 0 + + for op_idx in range(max_ops): + work_handles = [] + owned_send_tensors = [] + for peer_rank, ops in peer_ops_with_rank: + if op_idx >= len(ops): + continue + plan_op, p2p_op = ops[op_idx] + is_send = p2p_op.op is dist.isend or p2p_op.op is dist.send + stream = self._stream_pool[peer_to_stream_idx[peer_rank]] + with device_util.stream(stream): + # Prepare the payload on the same stream that consumes it. + # clone()/to() on the caller's default stream followed by + # isend() on this dedicated stream has no ordering edge; + # NCCL can otherwise read a partially written clone and + # silently deliver sparse NaN/Inf values. + tensor_for_transfer = ( + p2p_op.tensor.clone() if is_send else p2p_op.tensor + ) + if is_send: + # NCCL send/recv counts are expressed in elements of + # each side's dtype. A dtype mismatch therefore changes + # the wire size. Match the inference shard's dtype. + recv_dtype = getattr(plan_op.recv_shard_meta, "dtype", None) + if ( + recv_dtype is not None + and tensor_for_transfer.dtype != recv_dtype + ): + tensor_for_transfer = tensor_for_transfer.to(recv_dtype) + owned_send_tensors.append(tensor_for_transfer) + result = p2p_op.op( + tensor_for_transfer, + p2p_op.peer, + group=p2p_op.group, + ) + if p2p_op.op is dist.isend or p2p_op.op is dist.irecv: + work_handles.append(result) + total_ops += 1 + + for work in work_handles: + work.wait() + # ProcessGroupNCCL Work.wait() only guarantees that the CUDA work + # has been enqueued. The send clones must remain alive until NCCL + # has actually consumed them; otherwise the caching allocator can + # reuse their storage for the next batch and silently corrupt the + # transferred model. Drain this bounded batch before releasing it. + device_util.synchronize() + work_handles.clear() + owned_send_tensors.clear() + tensor_for_transfer = None + result = None + + return total_ops + + def _execute_ops_concurrent_packed( + self, + ops_dict, + peer_ranks, + max_pack_ops: int, + max_pack_bytes: int, + ) -> int: + """Execute one bounded expert pack per active peer and FIFO position.""" + from awex.util import device as device_util + + peer_batches_with_rank = [ + ( + peer_rank, + self._partition_expert_operations( + ops_dict[peer_rank], max_pack_ops, max_pack_bytes + ), + ) + for peer_rank in peer_ranks + if peer_rank in ops_dict + ] + if not peer_batches_with_rank: + return 0 + + peer_to_stream_idx = { + peer_rank: index % len(self._stream_pool) + for index, (peer_rank, _) in enumerate(peer_batches_with_rank) + } + max_batches = max(len(batches) for _, batches in peer_batches_with_rank) + logical_ops = 0 + wire_ops = 0 + packed_wire_ops = 0 + + for batch_idx in range(max_batches): + work_handles = [] + owned_send_tensors = [] + owned_recv_tensors = [] + pending_recv_unpacks = [] + for peer_rank, batches in peer_batches_with_rank: + if batch_idx >= len(batches): + continue + batch = batches[batch_idx] + plan_op, p2p_op = batch[0] + is_send = p2p_op.op is dist.isend or p2p_op.op is dist.send + is_recv = p2p_op.op is dist.irecv or p2p_op.op is dist.recv + stream = self._stream_pool[peer_to_stream_idx[peer_rank]] + with device_util.stream(stream): + if len(batch) == 1: + tensor_for_transfer = ( + p2p_op.tensor.clone() if is_send else p2p_op.tensor + ) + if is_send: + recv_dtype = self._operation_wire_dtype(plan_op, p2p_op) + if tensor_for_transfer.dtype != recv_dtype: + tensor_for_transfer = tensor_for_transfer.to(recv_dtype) + owned_send_tensors.append(tensor_for_transfer) + elif is_send: + tensor_for_transfer = self._pack_send_batch(batch) + owned_send_tensors.append(tensor_for_transfer) + elif is_recv: + recv_dtype = self._operation_wire_dtype(plan_op, p2p_op) + tensor_for_transfer = self._allocate_packed_recv_batch( + batch, recv_dtype + ) + owned_recv_tensors.append(tensor_for_transfer) + else: + raise RuntimeError( + "Expert packing only supports torch.distributed P2P ops" + ) + + result = p2p_op.op( + tensor_for_transfer, + p2p_op.peer, + group=p2p_op.group, + ) + if len(batch) > 1 and is_recv: + pending_recv_unpacks.append( + (stream, tensor_for_transfer, batch) + ) + + if p2p_op.op is dist.isend or p2p_op.op is dist.irecv: + work_handles.append((result, stream)) + logical_ops += len(batch) + wire_ops += 1 + packed_wire_ops += int(len(batch) > 1) + + for work, stream in work_handles: + # ProcessGroupNCCL wait establishes the completion dependency + # on the current stream. Use the same transfer stream that + # will consume a packed receive below. + with device_util.stream(stream): + work.wait() + for stream, packed, batch in pending_recv_unpacks: + with device_util.stream(stream): + self._unpack_recv_batch(packed, batch) + # Keep both packed send and receive buffers alive until their NCCL + # and foreach-copy work has drained from every active peer stream. + device_util.synchronize() + work_handles.clear() + owned_send_tensors.clear() + owned_recv_tensors.clear() + pending_recv_unpacks.clear() + tensor_for_transfer = None + result = None + + if getattr(self, "_expert_pack_stats", None) is not None: + self._expert_pack_stats["logical_ops"] += logical_ops + self._expert_pack_stats["wire_ops"] += wire_ops + self._expert_pack_stats["packed_wire_ops"] += packed_wire_ops + return wire_ops diff --git a/awex/util/cuda_ipc.py b/awex/util/cuda_ipc.py new file mode 100644 index 0000000..c6d96b7 --- /dev/null +++ b/awex/util/cuda_ipc.py @@ -0,0 +1,66 @@ +# Licensed to the Awex developers under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +"""Allocate portable CUDA IPC buffers without changing training allocations.""" + +from __future__ import annotations + +import re +from contextlib import contextmanager +from typing import Iterator + +import torch + + +@contextmanager +def cuda_ipc_allocation() -> Iterator[None]: + """Use ordinary CUDA storage for new IPC buffers, then restore the allocator. + + Expandable IPC handles are not portable across all PyTorch versions. Scope + this context to staging-buffer creation during serialized weight publishing, + while training is paused: the allocator switch is process-wide. Existing + training storage is untouched, and restoration precedes IPC serialization. + """ + settings = torch.cuda.memory._snapshot()["allocator_settings"] + if not settings["expandable_segments"]: + yield + return + + # Preserve runtime settings, not environment defaults. The setter resets + # other options (including split size and GC threshold) on every call. + original = settings["PYTORCH_CUDA_ALLOC_CONF"] + # A preceding setter can omit this option while retaining its True value. + # Replaying that string alone would leave our temporary False in effect. + restore = original + if not re.search(r"expandable_segments\s*:", original): + restore = ( + f"{original},expandable_segments:True" + if original + else "expandable_segments:True" + ) + staging = re.sub(r"expandable_segments\s*:\s*(True|False)", "", original) + staging = ",".join(part for part in staging.split(",") if part.strip()) + staging = ( + f"{staging},expandable_segments:False" + if staging + else "expandable_segments:False" + ) + try: + torch.cuda.memory._set_allocator_settings(staging) + yield + finally: + torch.cuda.memory._set_allocator_settings(restore) From c2360eaa4cd4fa3f96bd5384a3486b9c3d34c01f Mon Sep 17 00:00:00 2001 From: "chucai.dzq" Date: Tue, 22 Sep 2026 18:32:32 +0800 Subject: [PATCH 2/2] fix: bound strided transfer buffers and own Qwen4Exp converters --- awex/models/qwen4_exp.py | 359 +++++++++++++++++++++++++ awex/models/qwen4_exp_contract.py | 205 ++++++++++++++ awex/models/qwen4_exp_layout.py | 179 ++++++++++++ awex/tests/test_nccl_bounded_stream.py | 156 +++++------ awex/tests/test_qwen4_exp.py | 109 ++++++++ awex/transfer/nccl_bounded_stream.py | 71 +++-- 6 files changed, 968 insertions(+), 111 deletions(-) create mode 100644 awex/models/qwen4_exp.py create mode 100644 awex/models/qwen4_exp_contract.py create mode 100644 awex/models/qwen4_exp_layout.py create mode 100644 awex/tests/test_qwen4_exp.py diff --git a/awex/models/qwen4_exp.py b/awex/models/qwen4_exp.py new file mode 100644 index 0000000..92957ff --- /dev/null +++ b/awex/models/qwen4_exp.py @@ -0,0 +1,359 @@ +# Licensed to the Awex developers under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +"""Explicit Qwen4Exp AWEX mappings under development. + +These factories do not register themselves or enable the engine's AWEX guard. +PLE table/buffer residency requires a separate contract; +reject them instead of silently inheriting an unrelated model's conversion. +""" + +from __future__ import annotations + +import re +from collections.abc import Callable, Mapping +from functools import lru_cache + +import torch +from torch import nn + +from awex.models.qwen4_exp_contract import ( + Qwen4ExpFrozenContract, + mcore_visual_parameter_name, +) +from awex.models.qwen4_exp_layout import ( + Qwen4ExpGDNLayout, + pack_qwen4_exp_gated_qkv, +) + +_HC_WEIGHTS = ( + "hc_norm.weight", + "input_mix_weight_down.weight", + "input_mix_weight_up.weight", + "block_inject_weight.weight", +) +_REPLICATED_LAYER_WEIGHTS = frozenset( + f"{branch}.{weight}" + for branch in ("attn_hyper_connection", "mlp_hyper_connection") + for weight in _HC_WEIGHTS +) | frozenset( + f"ple.{weight}" + for weight in ( + "key_proj.weight", + "value_proj.weight", + "norm_key.weight", + "norm_query.weight", + "norm_conv.weight", + "conv1d.weight", + ) +) +_QSA_WEIGHTS = frozenset( + f"self_attn.indexer.{weight}" + for weight in ("index_qk_proj.weight", "q_layernorm.weight", "k_layernorm.weight") +) +_REPLICATED_LAYER_WEIGHTS |= _QSA_WEIGHTS +_MIXER_WEIGHTS = frozenset(f"hyper_connection_mixer.{w}" for w in _HC_WEIGHTS[:3]) + + +def _replicated_name(name: str) -> bool: + if name.startswith("model.") and name[len("model.") :] in _MIXER_WEIGHTS: + return True + match = re.fullmatch(r"model\.layers\.\d+\.(.+)", name) + return bool(match and match[1] in _REPLICATED_LAYER_WEIGHTS) + + +def _reject_pending_contract(name: str) -> None: + if ".ple_embedding." in name: + raise NotImplementedError( + f"Qwen4Exp AWEX requires an explicit table/buffer contract: {name}" + ) + if "hyper_connection" in name or ".ple." in name or ".indexer." in name: + if not _replicated_name(name): + raise NotImplementedError(f"Unknown Qwen4Exp replicated weight: {name}") + + +def _refresh_frozen_binding(converter, binder: Callable[[object], None] | None) -> None: + if binder is None: + raise ValueError("No Qwen4Exp frozen-contract binder was registered") + # Invalidate first: a failed refresh must never leave old Parameter references + # usable after model recovery or replacement. + for attribute in ( + "_qwen4_frozen_contract", + "_qwen4_original_parameters", + "_qwen4_local_table_names", + "_qwen4_local_visual_names", + "_qwen4_preserved_visual_names", + ): + converter.__dict__.pop(attribute, None) + try: + binder(converter) + if getattr(converter, "_qwen4_frozen_contract", None) is None: + raise ValueError("Qwen4Exp binder did not bind a frozen contract") + except Exception: + converter.__dict__.pop("_qwen4_frozen_contract", None) + raise + + +@lru_cache(maxsize=None) +def build_mcore_converter(binder: Callable[[object], None] | None = None): + from awex.converter.mcore_converter import _process_mcore_pp_name + from awex.models.qwen3_5 import _MCORE_CONVERTER_FACTORY + + base = _MCORE_CONVERTER_FACTORY() + + class McoreToHFWeightConverterQwen4Exp(base): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + if binder is not None: + self.refresh_frozen_contract() + + def refresh_frozen_contract(self) -> None: + _refresh_frozen_binding(self, binder) + + def bind_frozen_contract( + self, + contract: Qwen4ExpFrozenContract, + parameters: Mapping[str, nn.Parameter], + local_table_names: frozenset[str], + local_visual_names: frozenset[str] | None = None, + ) -> None: + contract.validate_actor_parameters( + parameters, local_table_names, local_visual_names + ) + self._qwen4_frozen_contract = contract + self._qwen4_original_parameters = parameters + self._qwen4_local_table_names = local_table_names + self._qwen4_local_visual_names = local_visual_names + + def _convert_attention_param(self, name, parameter, layer_number): + if name in ( + "self_attention.linear_qkv.weight", + "self_attention.linear_qkv.bias", + ): + cfg = self.hf_config + packed = pack_qwen4_exp_gated_qkv( + self._full_tp_tensor(parameter), + int(cfg.num_attention_heads), + int(cfg.num_key_value_heads), + int(cfg.head_dim), + int(self.infer_atten_tp_size), + ) + suffix = name.rsplit(".", 1)[1] + return [ + (f"self_attn.qkv_proj.{suffix}", self._take_train_tp_shard(packed)) + ] + if name == "self_attention.A_log": + # Native Qwen4Exp SGLang keeps A_log FP32; metadata and payload + # must describe the same converted dtype on the writer side. + return [("linear_attn.A_log", parameter.float())] + if name == "self_attention.out_norm.weight": + # FlashNext bridge uses ones-style GDN norm: no +1 offset. + return [("linear_attn.norm.weight", parameter)] + if name in ( + "self_attention.in_proj.weight", + "self_attention.conv1d.weight", + "self_attention.in_proj_qkvz.weight", + "self_attention.in_proj_ba.weight", + ): + cfg = self.hf_config + layout = Qwen4ExpGDNLayout( + cfg.linear_num_key_heads, + cfg.linear_num_value_heads, + cfg.linear_key_head_dim, + cfg.linear_value_head_dim, + ) + train_tp = int(self.rank_info.attn_tp_size) + infer_tp = int(self.infer_atten_tp_size) + full = self._full_tp_tensor(parameter) + if name in ( + "self_attention.in_proj_qkvz.weight", + "self_attention.in_proj_ba.weight", + ): + component = "qkvz" if "qkvz" in name else "ba" + packed = layout.pack_decoupled(full, train_tp, infer_tp, component) + return [ + ( + f"linear_attn.in_proj_{component}.weight", + self._take_train_tp_shard(packed), + ) + ] + if name == "self_attention.in_proj.weight": + qkvz, ba = layout.pack_input(full, train_tp, infer_tp) + return [ + ( + "linear_attn.in_proj_qkvz.weight", + self._take_train_tp_shard(qkvz), + ), + ( + "linear_attn.in_proj_ba.weight", + self._take_train_tp_shard(ba), + ), + ] + packed = layout.pack_conv(full, train_tp, infer_tp) + return [ + ("linear_attn.conv1d.weight", self._take_train_tp_shard(packed)) + ] + return super()._convert_attention_param(name, parameter, layer_number) + + @torch.no_grad() + def convert_param(self, name, parameter, vp_stage=None): + clean = name.replace("module.", "") + contract = getattr(self, "_qwen4_frozen_contract", None) + if clean.startswith("visual."): + if contract is None or contract.language_model_only: + raise ValueError( + "Actor visual conversion requires a bound vision contract" + ) + canonical = mcore_visual_parameter_name(name, contract) + contract.validate_actor_parameters( + self._qwen4_original_parameters, + self._qwen4_local_table_names, + self._qwen4_local_visual_names, + ) + if canonical not in self._qwen4_local_visual_names: + raise ValueError(f"Visual parameter is not owned locally: {name}") + return [] + if clean.startswith("language_model."): + clean = clean[len("language_model.") :] + if clean.startswith("decoder."): + global_name = _process_mcore_pp_name( + clean, + self.rank_info, + self.hf_config, + self.tf_config, + vp_stage=vp_stage, + pp_stage_layer_id_map=self._pp_stage_layer_id_map, + ) + canonical = "model." + global_name[len("decoder.") :] + canonical = canonical.replace( + ".self_attention.indexer.", ".self_attn.indexer." + ) + contract = getattr(self, "_qwen4_frozen_contract", None) + if contract is not None and contract.excludes(canonical, "actor"): + contract.validate_actor_parameters( + self._qwen4_original_parameters, + self._qwen4_local_table_names, + self._qwen4_local_visual_names, + ) + return [] + _reject_pending_contract(canonical) + if _replicated_name(canonical): + return [(canonical, parameter)] + if canonical == "model.final_layernorm.weight": + raise ValueError( + "Qwen4Exp has a final HC mixer, not a final layernorm" + ) + # Let the parent apply PP numbering once for inherited attention/MLP. + return super().convert_param(name, parameter, vp_stage=vp_stage) + + return McoreToHFWeightConverterQwen4Exp + + +@lru_cache(maxsize=None) +def build_sglang_converter(binder: Callable[[object], None] | None = None): + from awex.models.qwen3_5 import SGlangToHFWeightConverterQwen3_5 + + class SGlangToHFWeightConverterQwen4Exp(SGlangToHFWeightConverterQwen3_5): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + if binder is not None: + self.refresh_frozen_contract() + + def refresh_frozen_contract(self) -> None: + _refresh_frozen_binding(self, binder) + + def bind_frozen_contract( + self, + contract: Qwen4ExpFrozenContract, + parameters: Mapping[str, nn.Parameter], + preserved_visual_names: frozenset[str], + ) -> None: + contract.validate_inference_parameters(parameters, preserved_visual_names) + self._qwen4_frozen_contract = contract + self._qwen4_original_parameters = parameters + self._qwen4_preserved_visual_names = preserved_visual_names + + @torch.no_grad() + def convert_param(self, name, parameter): + canonical = name.replace("model.language_model.", "model.") + canonical = re.sub( + r"^(model\.layers\.\d+)\.indexer\.", r"\1.self_attn.indexer.", canonical + ) + if canonical.startswith("visual."): + canonical = "model." + canonical + contract = getattr(self, "_qwen4_frozen_contract", None) + if contract is not None and contract.excludes(canonical, "inference"): + contract.validate_inference_parameters( + self._qwen4_original_parameters, self._qwen4_preserved_visual_names + ) + return [] + _reject_pending_contract(canonical) + if _replicated_name(canonical): + return [(canonical, parameter)] + return super().convert_param(name, parameter) + + return SGlangToHFWeightConverterQwen4Exp + + +@lru_cache(maxsize=1) +def build_sharding_strategy(): + from awex.models.qwen3_5 import Qwen3_5ShardingStrategy + from awex.sharding.param_sharding import ShardingType + + class Qwen4ExpShardingStrategy(Qwen3_5ShardingStrategy): + def get_sharding_strategy(self, parameter_name, **kwargs): + _reject_pending_contract(parameter_name) + if _replicated_name(parameter_name): + return ShardingType.NO_SHARDING, 0, 1 + return super().get_sharding_strategy(parameter_name, **kwargs) + + return Qwen4ExpShardingStrategy + + +def register_qwen4_exp_awex( + *, + mcore_binder: Callable[[object], None] | None = None, + sglang_binder: Callable[[object], None] | None = None, +) -> None: + """Register explicit factories after AWEX has finished rebuilding its registry. + + Optional process-local, hashable binders run on every native construction, + including metadata resolvers and payload converters. They must obtain fresh + original Parameters and call bind_frozen_contract. Reuse the same callback + when registering again; call refresh_frozen_contract before each transfer + and after model recovery. Bindings are not transported between processes. + + Registration alone does not enable unsupported frozen-state exclusions or + remove the engine guard. Refuse to replace an unrelated upstream adapter. + """ + from awex.models.registry import ModelRegistry + + architecture = "Qwen4ExpForConditionalGeneration" + entry = { + "model_name": architecture, + "mcore_converter": build_mcore_converter + if mcore_binder is None + else build_mcore_converter(mcore_binder), + "sglang_converter": build_sglang_converter + if sglang_binder is None + else build_sglang_converter(sglang_binder), + "sharding_strategy": build_sharding_strategy(), + } + existing = ModelRegistry.models.get(architecture) + if existing is not None and existing != entry: + raise ValueError("A different Qwen4Exp AWEX adapter is already registered") + ModelRegistry.models[architecture] = entry diff --git a/awex/models/qwen4_exp_contract.py b/awex/models/qwen4_exp_contract.py new file mode 100644 index 0000000..b631f0d --- /dev/null +++ b/awex/models/qwen4_exp_contract.py @@ -0,0 +1,205 @@ +# Licensed to the Awex developers under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +"""Exact metadata boundaries for Qwen4Exp immutable inference state. + +This declaration does not prove tensor values or lifecycle preservation. Runtime +integration must additionally bind the checkpoint evidence and visual backup, +validate original parameters before each exchange, and use the same declaration +for metadata and payload converters. No automatic registration or exclusion is +performed here. +""" + +from __future__ import annotations + +import re +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any, Literal + +import torch +from torch import nn + + +def mcore_visual_parameter_name(name: str, contract: Qwen4ExpFrozenContract) -> str: + """Map the replicated ModelScope HF vision tower to receiver identities.""" + while name.startswith("module."): + name = name[len("module.") :] + if not name.startswith("visual.visual."): + raise ValueError(f"Unsupported Qwen4Exp actor visual parameter: {name}") + canonical = "model.visual." + name[len("visual.visual.") :] + if canonical not in contract.visual_parameter_names: + canonical = canonical.replace(".attn.qkv.", ".attn.qkv_proj.") + if canonical not in contract.visual_parameter_names: + raise ValueError(f"Actor visual parameter is outside frozen contract: {name}") + return canonical + + +@dataclass(frozen=True) +class Qwen4ExpFrozenContract: + checkpoint_manifest_sha256: str + ple_table_names: frozenset[str] + visual_parameter_names: frozenset[str] + language_model_only: bool + freeze_ple_table: bool + schema_version: int = 1 + + def __post_init__(self) -> None: + if type(self.schema_version) is not int or self.schema_version not in (1, 2): + raise ValueError("Unsupported Qwen4Exp frozen contract schema") + if ( + type(self.language_model_only) is not bool + or self.freeze_ple_table is not True + ): + raise ValueError( + "Frozen exclusions require an explicit model mode and frozen PLE" + ) + if self.schema_version == 1 and not self.language_model_only: + raise ValueError("Vision actors require frozen contract schema 2") + if not re.fullmatch(r"[0-9a-f]{64}", self.checkpoint_manifest_sha256): + raise ValueError("Expected a SHA256 checkpoint manifest identity") + for names in (self.ple_table_names, self.visual_parameter_names): + if not isinstance(names, frozenset) or not names: + raise ValueError("Frozen parameter names must be nonempty frozen sets") + for name in self.ple_table_names: + if not re.fullmatch( + r"model\.layers\.\d+\.ple\.ple_embedding\.ngram_embedding\.weight", + name, + ): + raise ValueError(f"Invalid frozen PLE table name: {name}") + for name in self.visual_parameter_names: + if not re.fullmatch( + r"model\.visual\.[A-Za-z0-9_]+(?:\.[A-Za-z0-9_]+)*", name + ): + raise ValueError(f"Invalid frozen visual parameter name: {name}") + + def to_dict(self) -> dict[str, Any]: + return { + "schema_version": self.schema_version, + "checkpoint_manifest_sha256": self.checkpoint_manifest_sha256, + "language_model_only": self.language_model_only, + "freeze_ple_table": self.freeze_ple_table, + "ple_table_names": sorted(self.ple_table_names), + "visual_parameter_names": sorted(self.visual_parameter_names), + } + + @classmethod + def from_dict(cls, payload: Mapping[str, Any]) -> Qwen4ExpFrozenContract: + data = dict(payload) + expected = { + "schema_version", + "checkpoint_manifest_sha256", + "language_model_only", + "freeze_ple_table", + "ple_table_names", + "visual_parameter_names", + } + if data.keys() != expected: + raise ValueError("Missing or unexpected Qwen4Exp frozen contract fields") + for key in ("ple_table_names", "visual_parameter_names"): + names = data[key] + if not isinstance(names, list) or not all( + isinstance(n, str) for n in names + ): + raise ValueError(f"Expected a string list for {key}") + if len(names) != len(set(names)): + raise ValueError(f"Duplicate names in {key}") + data[key] = frozenset(names) + return cls(**data) + + def excludes(self, name: str, side: Literal["actor", "inference"]) -> bool: + if side not in ("actor", "inference"): + raise ValueError(f"Unknown contract side: {side}") + return name in self.ple_table_names or ( + (side == "inference" or not self.language_model_only) + and name in self.visual_parameter_names + ) + + @staticmethod + def _validate_table(name: str, parameter: nn.Parameter) -> None: + if not isinstance(parameter, nn.Parameter): + raise TypeError( + f"Validate original Parameter objects, not detached tensors: {name}" + ) + if parameter.requires_grad: + raise ValueError(f"PLE table is trainable: {name}") + if parameter.dtype != torch.bfloat16 or parameter.ndim != 2: + raise ValueError(f"Expected the validated BF16 PLE table layout: {name}") + + def validate_actor_parameters( + self, + parameters: Mapping[str, nn.Parameter], + local_table_names: frozenset[str], + local_visual_names: frozenset[str] | None = None, + ) -> None: + """Validate canonical original parameters on this PP stage, before detach. + + The caller must verify global PP ownership coverage separately; a PP stage + without PLE legitimately has an empty local table set. + """ + if not local_table_names <= self.ple_table_names: + raise ValueError("Local PLE ownership is outside the frozen contract") + observed_visual = frozenset( + n for n in parameters if n.startswith("model.visual.") + ) + if self.language_model_only and observed_visual: + raise ValueError( + "Language-only actor unexpectedly contains visual parameters" + ) + if not self.language_model_only: + if local_visual_names not in (frozenset(), self.visual_parameter_names): + raise ValueError("Explicit complete local visual ownership is required") + if observed_visual != local_visual_names: + raise ValueError("Actor visual parameters do not match local ownership") + for name in observed_visual: + parameter = parameters[name] + if not isinstance(parameter, nn.Parameter): + raise TypeError( + f"Validate original visual Parameter, not detached tensor: {name}" + ) + if parameter.requires_grad: + raise ValueError(f"Frozen visual parameter is trainable: {name}") + if parameter.dtype != torch.bfloat16: + raise ValueError(f"Expected BF16 frozen visual parameter: {name}") + observed = {name for name in parameters if ".ple_embedding." in name} + if observed != local_table_names: + raise ValueError( + "Actor PLE parameters do not match declared local ownership" + ) + for name in local_table_names: + self._validate_table(name, parameters[name]) + + def validate_inference_parameters( + self, + parameters: Mapping[str, nn.Parameter], + preserved_visual_names: frozenset[str], + ) -> None: + """Require exact exclusions and the same keys in the visual backup path.""" + observed_visual = {n for n in parameters if n.startswith("model.visual.")} + if observed_visual != self.visual_parameter_names: + raise ValueError( + "Inference visual parameters differ from the frozen contract" + ) + if preserved_visual_names != self.visual_parameter_names: + raise ValueError("Visual preservation keys differ from transfer exclusions") + observed_tables = {n for n in parameters if ".ple_embedding." in n} + if observed_tables != self.ple_table_names: + raise ValueError("Inference PLE parameters differ from the frozen contract") + for name in self.ple_table_names: + self._validate_table(name, parameters[name]) + if parameters[name].device.type != "cpu": + raise ValueError(f"Expected the validated CPU PLE residency: {name}") diff --git a/awex/models/qwen4_exp_layout.py b/awex/models/qwen4_exp_layout.py new file mode 100644 index 0000000..e5e50ca --- /dev/null +++ b/awex/models/qwen4_exp_layout.py @@ -0,0 +1,179 @@ +# Licensed to the Awex developers under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +"""Qwen4Exp bridge GDN rows to AWEX's inference-rank-major representation. + +Inputs are full tensors gathered along dimension zero in training TP rank +order. They contain whole key-head groups, not rank-local Q/K/V categories. +Outputs retain the full dimension zero; the AWEX writer subsequently takes +its training-rank slice and the transfer plan redistributes inference slices. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Literal + +import torch + + +@dataclass(frozen=True) +class Qwen4ExpGDNLayout: + num_key_heads: int + num_value_heads: int + key_head_dim: int + value_head_dim: int + + def __post_init__(self) -> None: + for value in ( + self.num_key_heads, + self.num_value_heads, + self.key_head_dim, + self.value_head_dim, + ): + if value <= 0: + raise ValueError("GDN head counts and dimensions must be positive") + if self.num_value_heads % self.num_key_heads: + raise ValueError("GDN value heads must divide into whole key-head groups") + + def _categories( + self, + parameter: torch.Tensor, + train_tp_size: int, + infer_tp_size: int, + *, + component: Literal["input", "conv", "qkvz", "ba"], + ) -> tuple[torch.Tensor, ...]: + for size in (train_tp_size, infer_tp_size): + if size <= 0 or self.num_key_heads % size: + raise ValueError("GDN TP sizes must divide the key-head count") + ratio = self.num_value_heads // self.num_key_heads + value_width = ratio * self.value_head_dim + qkv = (self.key_head_dim, self.key_head_dim, value_width) + dimensions = { + "input": (*qkv, value_width, ratio, ratio), + "conv": qkv, + "qkvz": (*qkv, value_width), + "ba": (ratio, ratio), + } + if component not in dimensions: + raise ValueError(f"Unknown GDN component: {component}") + widths = dimensions[component] + expected_rows = self.num_key_heads * sum(widths) + if parameter.ndim < 2 or parameter.shape[0] != expected_rows: + raise ValueError( + f"Expected full GDN tensor with {expected_rows} rows; " + f"got {tuple(parameter.shape)}" + ) + grouped = parameter.reshape( + self.num_key_heads, sum(widths), *parameter.shape[1:] + ) + return tuple( + part.reshape(-1, *parameter.shape[1:]) + for part in grouped.split(widths, dim=1) + ) + + @staticmethod + def _pack(categories: tuple[torch.Tensor, ...], infer_tp_size: int) -> torch.Tensor: + shards = [category.chunk(infer_tp_size, dim=0) for category in categories] + return torch.cat( + [ + torch.cat([parts[rank] for parts in shards], dim=0) + for rank in range(infer_tp_size) + ], + dim=0, + ).contiguous() + + def pack_input( + self, + parameter: torch.Tensor, + train_tp_size: int, + infer_tp_size: int, + ) -> tuple[torch.Tensor, torch.Tensor]: + categories = self._categories( + parameter, train_tp_size, infer_tp_size, component="input" + ) + return ( + self._pack(categories[:4], infer_tp_size), + self._pack(categories[4:], infer_tp_size), + ) + + def pack_conv( + self, + parameter: torch.Tensor, + train_tp_size: int, + infer_tp_size: int, + ) -> torch.Tensor: + categories = self._categories( + parameter, train_tp_size, infer_tp_size, component="conv" + ) + return self._pack(categories, infer_tp_size) + + def pack_decoupled( + self, + parameter: torch.Tensor, + train_tp_size: int, + infer_tp_size: int, + component: Literal["qkvz", "ba"], + ) -> torch.Tensor: + if component not in ("qkvz", "ba"): + raise ValueError(f"Expected a decoupled GDN component, got {component}") + categories = self._categories( + parameter, train_tp_size, infer_tp_size, component=component + ) + return self._pack(categories, infer_tp_size) + + +def pack_qwen4_exp_gated_qkv( + parameter: torch.Tensor, + num_heads: int, + num_kv_heads: int, + head_dim: int, + infer_tp_size: int, +) -> torch.Tensor: + """Preserve bridge Q/gate head interleaving and replicate inference KV heads.""" + if min(num_heads, num_kv_heads, head_dim, infer_tp_size) <= 0: + raise ValueError("Attention geometry and TP size must be positive") + if num_heads % num_kv_heads or num_heads % infer_tp_size: + raise ValueError("Query heads must divide into KV groups and TP shards") + if max(num_kv_heads, infer_tp_size) % min(num_kv_heads, infer_tp_size): + raise ValueError("KV heads and TP size must divide one another") + query_rows = 2 * (num_heads // num_kv_heads) * head_dim + rows = query_rows + 2 * head_dim + if parameter.ndim < 1 or parameter.shape[0] != num_kv_heads * rows: + raise ValueError("Unexpected full Qwen4Exp gated QKV tensor shape") + tail = parameter.shape[1:] + groups = parameter.reshape(num_kv_heads, rows, *tail) + # Bridge directly concatenates HF q_proj (already Q/gate interleaved), K, V. + query = groups[:, :query_rows].reshape(num_heads, 2 * head_dim, *tail) + key = groups[:, query_rows : query_rows + head_dim] + value = groups[:, query_rows + head_dim :] + query_parts = query.chunk(infer_tp_size, dim=0) + if infer_tp_size >= num_kv_heads: + replicas = infer_tp_size // num_kv_heads + key_parts = [key[rank // replicas] for rank in range(infer_tp_size)] + value_parts = [value[rank // replicas] for rank in range(infer_tp_size)] + else: + key_parts = key.chunk(infer_tp_size, dim=0) + value_parts = value.chunk(infer_tp_size, dim=0) + return torch.cat( + [ + torch.cat([part.reshape(-1, *tail) for part in parts], dim=0) + for parts in zip(query_parts, key_parts, value_parts) + ], + dim=0, + ).contiguous() diff --git a/awex/tests/test_nccl_bounded_stream.py b/awex/tests/test_nccl_bounded_stream.py index 7c8046e..864b1d4 100644 --- a/awex/tests/test_nccl_bounded_stream.py +++ b/awex/tests/test_nccl_bounded_stream.py @@ -28,79 +28,85 @@ ) -@pytest.mark.parametrize("pending_slice_copy", [False, True]) -def test_bounded_transport_defers_send_clones_until_execution( - monkeypatch, pending_slice_copy +@pytest.mark.parametrize("pack_ops", [1, 64]) +def test_bounded_transport_defers_strided_buffers_until_execution( + monkeypatch, pack_ops ): - """Building an AWEX transfer plan retains views instead of model-sized clones.""" + """Real column slices remain views; wire buffers live for just one batch.""" + import weakref + from awex.transfer import nccl_stream_batch from awex.util import device as device_util - class _SourceTensor: - def __init__(self) -> None: - self.clone_calls = 0 - - def clone(self): - self.clone_calls += 1 - return self - - source = _SourceTensor() - source.ready = not pending_slice_copy - send_op = SimpleNamespace( - send_shard_meta=SimpleNamespace(name="weight"), - recv_rank=1, - ) - send_plan = SimpleNamespace(operations={1: [send_op]}) - recv_op = SimpleNamespace(recv_shard_meta=SimpleNamespace(name="weight")) - recv_plan = SimpleNamespace(operations={1: [recv_op]}) - recv_storage = torch.full((2, 4), float("nan")) - recv_target = recv_storage[:, ::2] - expected_recv = torch.arange(4, dtype=torch.float32).reshape(2, 2) - transport = object.__new__(BoundedMemoryNcclColocateStreamBatchTransport) - - def _inspect_plan( - transfer_rank, - world_size, - all_send_p2p_ops, - all_recv_p2p_ops, - weights_update_group, - rank_coordinate, - step_id, - ) -> None: - del ( - transfer_rank, - world_size, - weights_update_group, - rank_coordinate, - step_id, + sources = {f"w{i}": torch.arange(32).reshape(4, 8).float() + i for i in range(3)} + targets = {name: torch.full((4, 8), float("nan")) for name in sources} + plan_ops = [ + SimpleNamespace( + send_shard_meta=SimpleNamespace(name=name), + recv_shard_meta=SimpleNamespace(name=name, dtype=torch.float32), + train_slices=(slice(None), slice(2, 6)), + inf_slices=(slice(None), slice(None, None, 2)), + recv_rank=1, + overlap_shape=(4, 4), + param_class="expert", ) - assert all_send_p2p_ops[1][0][1].tensor is source - assert source.clone_calls == 0 - # A materialized slice cannot be consumed on the transfer stream until - # its asynchronous producer on the caller stream has completed. - assert source.ready - recv_buffer = all_recv_p2p_ops[1][0][1].tensor - assert recv_buffer.is_contiguous() - assert recv_buffer.data_ptr() != recv_target.data_ptr() - recv_buffer.copy_(expected_recv) + for name in sources + ] + transport = object.__new__(BoundedMemoryNcclColocateStreamBatchTransport) + transport._stream_pool = [object()] + live = [] + payloads = [] + recv_index = 0 + + class Work: + def wait(self): + pass + + def send(tensor, peer, group): + assert tensor.is_contiguous() + assert all(ref() is None for ref in live) + live[:] = [weakref.ref(tensor)] + payloads.append(tensor.clone()) + return Work() + + def recv(tensor, peer, group): + nonlocal recv_index + assert tensor.is_contiguous() + assert all(ref() is None for ref in live) + live[:] = [weakref.ref(tensor)] + tensor.copy_(payloads[recv_index]) + recv_index += 1 + return Work() + + def execute(rank, world, sends, recvs, group, coordinate, step): + # No planning-stage contiguous copies, on either side. + for op, p2p in sends[1]: + assert not p2p.tensor.is_contiguous() + assert ( + p2p.tensor.untyped_storage().data_ptr() + == sources[op.send_shard_meta.name].untyped_storage().data_ptr() + ) + for op, p2p in recvs[1]: + assert not p2p.tensor.is_contiguous() + assert ( + p2p.tensor.untyped_storage().data_ptr() + == targets[op.recv_shard_meta.name].untyped_storage().data_ptr() + ) + transport._execute_ops_concurrent(sends, [1]) + assert all(ref() is None for ref in live) + transport._execute_ops_concurrent(recvs, [1]) + assert all(ref() is None for ref in live) + monkeypatch.setenv("AWEX_EXPERT_PACK_OPS", str(pack_ops)) monkeypatch.setattr(transport, "_validate_pack_config", lambda group: None) - transport.execute_recursive_partition_stream_transfer = _inspect_plan - monkeypatch.setattr( - nccl_stream_batch, - "hang_detector", - SimpleNamespace(submit=lambda *args, **kwargs: None), - ) + transport.execute_recursive_partition_stream_transfer = execute monkeypatch.setattr( - "awex.transfer.nccl_comm.validate_rank_mappings", lambda *args: None - ) - monkeypatch.setattr( - "awex.transfer.transfer_plan.slice_tensor", - lambda tensor, *args, **kwargs: tensor, - ) - monkeypatch.setattr( - device_util, "synchronize", lambda: setattr(source, "ready", True) + nccl_stream_batch, "hang_detector", SimpleNamespace(submit=lambda *a, **k: None) ) + monkeypatch.setattr(device_util, "stream", lambda stream: nullcontext()) + monkeypatch.setattr(device_util, "synchronize", lambda: None) + monkeypatch.setattr(torch.distributed, "isend", send) + monkeypatch.setattr(torch.distributed, "irecv", recv) monkeypatch.setattr( torch.distributed, "P2POp", @@ -108,24 +114,24 @@ def _inspect_plan( op=op, tensor=tensor, peer=peer, group=group ), ) - transport.update_weights_in_colocate_mode( train_to_infer_device_mapping={0: 0, 1: 1}, infer_to_train_device_mapping={0: 0, 1: 1}, transfer_rank=0, rank_coordinate="0-0-0", world_size=2, - send_transfer_plan=send_plan, - recv_transfer_plan=recv_plan, + send_transfer_plan=SimpleNamespace(operations={1: plan_ops}), + recv_transfer_plan=SimpleNamespace(operations={1: plan_ops}), weights_update_group=object(), - send_parameters={"weight": source}, - recv_parameters={"weight": recv_target}, + send_parameters=sources, + recv_parameters=targets, step_id=1, ) - - assert source.clone_calls == 0 - torch.testing.assert_close(recv_target, expected_recv, rtol=0, atol=0) - assert torch.isnan(recv_storage[:, 1::2]).all() + for name in sources: + torch.testing.assert_close( + targets[name][:, ::2], sources[name][:, 2:6], rtol=0, atol=0 + ) + assert torch.isnan(targets[name][:, 1::2]).all() def test_bounded_transport_releases_each_send_clone_batch(monkeypatch): @@ -143,7 +149,7 @@ def __del__(self) -> None: counters["live"] -= 1 class _SourceTensor: - def clone(self): + def clone(self, *, memory_format): counters["clones"] += 1 return _Clone() @@ -257,7 +263,7 @@ def __exit__(self, exc_type, exc_value, traceback): class _SourceTensor: dtype = torch.bfloat16 - def clone(self): + def clone(self, *, memory_format): assert state["active_stream"] is transfer_stream return self diff --git a/awex/tests/test_qwen4_exp.py b/awex/tests/test_qwen4_exp.py new file mode 100644 index 0000000..e20da68 --- /dev/null +++ b/awex/tests/test_qwen4_exp.py @@ -0,0 +1,109 @@ +# Licensed to the Awex developers under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +import pytest +import torch + +from awex.models.qwen4_exp_layout import Qwen4ExpGDNLayout + + +def _labels(heads, widths, tail): + # Encode each semantic coordinate independently of the implementation's + # reshape/split operations. MCore concatenates whole head groups. + rows = [] + lookup = {} + for head in range(heads): + for category, width in enumerate(widths): + for channel in range(width): + value = category * 100000 + head * 1000 + channel + rows.append(value) + lookup[category, head, channel] = value + tensor = torch.tensor(rows, dtype=torch.int64) + return tensor.reshape(-1, *([1] * len(tail))).expand(-1, *tail).clone(), lookup + + +def _expected(lookup, heads, widths, categories, infer_tp, tail): + rows = [] + for rank in range(infer_tp): + for category in categories: + for head in range(rank * heads // infer_tp, (rank + 1) * heads // infer_tp): + for channel in range(widths[category]): + rows.append(lookup[category, head, channel]) + return ( + torch.tensor(rows, dtype=torch.int64) + .reshape(-1, *([1] * len(tail))) + .expand(-1, *tail) + ) + + +@pytest.mark.parametrize("train_tp", [1, 2, 4, 8]) +@pytest.mark.parametrize("infer_tp", [1, 2, 4, 8]) +def test_gdn_packing_multiple_heads_preserves_semantic_coordinates(train_tp, infer_tp): + """Actual model head geometry; narrow hidden width keeps this CPU test small.""" + layout = Qwen4ExpGDNLayout(16, 48, 128, 128) + widths = (128, 128, 384, 384, 3, 3) + source, lookup = _labels(16, widths, (3,)) + original = source.clone() + qkvz, ba = layout.pack_input(source, train_tp, infer_tp) + for actual, categories in ((qkvz, range(4)), (ba, range(4, 6))): + expected = _expected(lookup, 16, widths, categories, infer_tp, (3,)) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + torch.testing.assert_close(source, original, rtol=0, atol=0) + + conv, conv_lookup = _labels(16, widths[:3], (1, 4)) + actual_conv = layout.pack_conv(conv, train_tp, infer_tp) + expected_conv = _expected(conv_lookup, 16, widths[:3], range(3), infer_tp, (1, 4)) + torch.testing.assert_close(actual_conv, expected_conv, rtol=0, atol=0) + for component, sizes in (("qkvz", widths[:4]), ("ba", widths[4:])): + decoupled, labels = _labels(16, sizes, (3,)) + actual = layout.pack_decoupled(decoupled, train_tp, infer_tp, component) + expected = _expected(labels, 16, sizes, range(len(sizes)), infer_tp, (3,)) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +@pytest.mark.parametrize("infer_tp", [1, 2, 4, 8]) +@pytest.mark.parametrize("tail", [(), (3,)]) +def test_gated_qkv_preserves_head_gate_pairs_and_replicates_kv(infer_tp, tail): + from awex.models.qwen4_exp_layout import pack_qwen4_exp_gated_qkv + + heads, kv_heads, dim = 24, 2, 4 + queries = [ + [10000 + h * 100 + g * 10 + c for g in range(2) for c in range(dim)] + for h in range(heads) + ] + keys = [[20000 + h * 100 + c for c in range(dim)] for h in range(kv_heads)] + values = [[30000 + h * 100 + c for c in range(dim)] for h in range(kv_heads)] + source = [] + for kv in range(kv_heads): + for head in range(kv * 12, (kv + 1) * 12): + source.extend(queries[head]) + source.extend(keys[kv]) + source.extend(values[kv]) + expected = [] + for rank in range(infer_tp): + for head in range(rank * heads // infer_tp, (rank + 1) * heads // infer_tp): + expected.extend(queries[head]) + owners = range(kv_heads) if infer_tp == 1 else [rank // (infer_tp // kv_heads)] + for category in (keys, values): + for owner in owners: + expected.extend(category[owner]) + + def tensor(rows): + return torch.tensor(rows).reshape(-1, *([1] * len(tail))).expand(-1, *tail) + + actual = pack_qwen4_exp_gated_qkv(tensor(source), heads, kv_heads, dim, infer_tp) + torch.testing.assert_close(actual, tensor(expected), rtol=0, atol=0) diff --git a/awex/transfer/nccl_bounded_stream.py b/awex/transfer/nccl_bounded_stream.py index 76966bc..8fd65e8 100644 --- a/awex/transfer/nccl_bounded_stream.py +++ b/awex/transfer/nccl_bounded_stream.py @@ -119,7 +119,6 @@ def update_weights_in_colocate_mode( validate_rank_mappings, ) from awex.transfer.nccl_stream_batch import hang_detector - from awex.transfer.transfer_plan import slice_tensor from awex.util import device as device_util logger.info( @@ -163,32 +162,20 @@ def update_weights_in_colocate_mode( all_send_p2p_ops = {} all_recv_p2p_ops = {} tensors_to_copy = [] - train_slice_context = {} - non_contiguous_tensor_pairs = [] for peer_rank, ops in send_ops.items(): mapped_peer_rank = train_to_infer_device_mapping.get(peer_rank, peer_rank) if mapped_peer_rank == transfer_rank: for op in ops: send_tensor = send_parameters[op.send_shard_meta.name] - tensor_sliced = slice_tensor( - send_tensor, - op, - True, - slice_context=train_slice_context, - ) + tensor_sliced = send_tensor[op.train_slices] tensors_to_copy.append(tensor_sliced) continue p2p_ops = [] for op in ops: send_tensor = send_parameters[op.send_shard_meta.name] - tensor_sliced = slice_tensor( - send_tensor, - op, - True, - slice_context=train_slice_context, - ) + tensor_sliced = send_tensor[op.train_slices] recv_rank = train_to_infer_device_mapping.get( op.recv_rank, op.recv_rank ) @@ -210,13 +197,7 @@ def update_weights_in_colocate_mode( p2p_ops = [] for op in ops: recv_tensor = recv_parameters[op.recv_shard_meta.name] - tensor_sliced = slice_tensor(recv_tensor, op, False) - if not tensor_sliced.is_contiguous(): - original_tensor = tensor_sliced - tensor_sliced = torch.empty_like( - tensor_sliced, memory_format=torch.contiguous_format - ) - non_contiguous_tensor_pairs.append((original_tensor, tensor_sliced)) + tensor_sliced = recv_tensor[op.inf_slices] p2p_op = dist.P2POp( dist.irecv if async_op else dist.recv, tensor_sliced, @@ -237,9 +218,8 @@ def update_weights_in_colocate_mode( else: logger.info("No tensors to copy for %s", task_id) - # slice_tensor may materialize send slices on the caller stream. - # Finish planning copies before independent transfer streams consume - # them, including ranks with no local copy to synchronize implicitly. + # Finish source writes and local copies before transfer streams read + # source views, including ranks with no local copy to synchronize. device_util.synchronize() future = Future() @@ -260,11 +240,6 @@ def update_weights_in_colocate_mode( rank_coordinate, step_id, ) - if non_contiguous_tensor_pairs: - with torch.no_grad(): - for original_tensor, recv_tensor in non_contiguous_tensor_pairs: - original_tensor.copy_(recv_tensor) - non_contiguous_tensor_pairs.clear() device_util.synchronize() future.set_result(True) if self._expert_pack_stats is not None: @@ -432,6 +407,7 @@ def _execute_ops_concurrent_unpacked(self, ops_dict, peer_ranks): for op_idx in range(max_ops): work_handles = [] owned_send_tensors = [] + pending_recv_copies = [] for peer_rank, ops in peer_ops_with_rank: if op_idx >= len(ops): continue @@ -445,7 +421,9 @@ def _execute_ops_concurrent_unpacked(self, ops_dict, peer_ranks): # NCCL can otherwise read a partially written clone and # silently deliver sparse NaN/Inf values. tensor_for_transfer = ( - p2p_op.tensor.clone() if is_send else p2p_op.tensor + p2p_op.tensor.clone(memory_format=torch.contiguous_format) + if is_send + else p2p_op.tensor ) if is_send: # NCCL send/recv counts are expressed in elements of @@ -458,17 +436,28 @@ def _execute_ops_concurrent_unpacked(self, ops_dict, peer_ranks): ): tensor_for_transfer = tensor_for_transfer.to(recv_dtype) owned_send_tensors.append(tensor_for_transfer) + elif not p2p_op.tensor.is_contiguous(): + tensor_for_transfer = torch.empty_like( + p2p_op.tensor, memory_format=torch.contiguous_format + ) + pending_recv_copies.append( + (stream, p2p_op.tensor, tensor_for_transfer) + ) result = p2p_op.op( tensor_for_transfer, p2p_op.peer, group=p2p_op.group, ) if p2p_op.op is dist.isend or p2p_op.op is dist.irecv: - work_handles.append(result) + work_handles.append((result, stream)) total_ops += 1 - for work in work_handles: - work.wait() + for work, stream in work_handles: + with device_util.stream(stream): + work.wait() + for stream, destination, received in pending_recv_copies: + with device_util.stream(stream), torch.no_grad(): + destination.copy_(received) # ProcessGroupNCCL Work.wait() only guarantees that the CUDA work # has been enqueued. The send clones must remain alive until NCCL # has actually consumed them; otherwise the caching allocator can @@ -477,6 +466,8 @@ def _execute_ops_concurrent_unpacked(self, ops_dict, peer_ranks): device_util.synchronize() work_handles.clear() owned_send_tensors.clear() + pending_recv_copies.clear() + destination = received = None tensor_for_transfer = None result = None @@ -530,13 +521,20 @@ def _execute_ops_concurrent_packed( with device_util.stream(stream): if len(batch) == 1: tensor_for_transfer = ( - p2p_op.tensor.clone() if is_send else p2p_op.tensor + p2p_op.tensor.clone(memory_format=torch.contiguous_format) + if is_send + else p2p_op.tensor ) if is_send: recv_dtype = self._operation_wire_dtype(plan_op, p2p_op) if tensor_for_transfer.dtype != recv_dtype: tensor_for_transfer = tensor_for_transfer.to(recv_dtype) owned_send_tensors.append(tensor_for_transfer) + elif not p2p_op.tensor.is_contiguous(): + tensor_for_transfer = self._allocate_packed_recv_batch( + batch, self._operation_wire_dtype(plan_op, p2p_op) + ) + owned_recv_tensors.append(tensor_for_transfer) elif is_send: tensor_for_transfer = self._pack_send_batch(batch) owned_send_tensors.append(tensor_for_transfer) @@ -556,7 +554,7 @@ def _execute_ops_concurrent_packed( p2p_op.peer, group=p2p_op.group, ) - if len(batch) > 1 and is_recv: + if is_recv and tensor_for_transfer is not p2p_op.tensor: pending_recv_unpacks.append( (stream, tensor_for_transfer, batch) ) @@ -583,6 +581,7 @@ def _execute_ops_concurrent_packed( owned_send_tensors.clear() owned_recv_tensors.clear() pending_recv_unpacks.clear() + packed = None tensor_for_transfer = None result = None