From 6710dce8353ebc56fbdb52e97d3b7590aa463e16 Mon Sep 17 00:00:00 2001 From: jinyx5 Date: Tue, 1 Sep 2026 00:45:14 +0800 Subject: [PATCH 1/7] fix(models): gate channels_last_3d activations to CUDA devices The VAE weight-side conversion to channels_last_3d is already guarded by cuDNN availability, but the activation-side casts in CausalConv3d.forward and the spatial-parallel halo conv were unconditional. NPU and CPU reject channels_last_3d activations, so Wan2.2 VAE forward failed on Ascend with ERR01007 ("NPU contiguous operator only supported contiguous memory format"). Route non-CUDA tensors through standard contiguous instead. Verified: Wan2.2-TI2V-5B text-to-video smoke on Ascend 910B2 (CANN 8.2, torch_npu 2.9.0) completes end to end; CUDA path unchanged. --- telefuser/distributed/vae_spatial.py | 6 +++++- telefuser/models/wan_video_vae.py | 6 +++++- 2 files changed, 10 insertions(+), 2 deletions(-) diff --git a/telefuser/distributed/vae_spatial.py b/telefuser/distributed/vae_spatial.py index 85f2f410..ebeb30fe 100644 --- a/telefuser/distributed/vae_spatial.py +++ b/telefuser/distributed/vae_spatial.py @@ -128,7 +128,11 @@ def _spatial_causal_conv3d_forward( if any(padding): tensor = F.pad(tensor, padding) tensor = _exchange_height_halo(module, tensor, module._height_halo_size) - tensor = tensor.contiguous(memory_format=torch.channels_last_3d) + if tensor.device.type == "cuda": + tensor = tensor.contiguous(memory_format=torch.channels_last_3d) + else: + # channels_last_3d activations are only supported by cuDNN; NPU/CPU require standard contiguous. + tensor = tensor.contiguous() return F.conv3d( tensor, module.weight, diff --git a/telefuser/models/wan_video_vae.py b/telefuser/models/wan_video_vae.py index 22b4bc7c..941f8056 100644 --- a/telefuser/models/wan_video_vae.py +++ b/telefuser/models/wan_video_vae.py @@ -119,7 +119,11 @@ def forward(self, x: torch.Tensor, cache_x: torch.Tensor | None = None) -> torch x = torch.cat([cache_x, x], dim=2) padding[4] -= cache_x.shape[2] x = F.pad(x, padding) - x = x.contiguous(memory_format=torch.channels_last_3d) + if x.device.type == "cuda": + x = x.contiguous(memory_format=torch.channels_last_3d) + else: + # channels_last_3d activations are only supported by cuDNN; NPU/CPU require standard contiguous. + x = x.contiguous() return super().forward(x) From ba34f29cedf814746fb066e4473da7a39f72e9ba Mon Sep 17 00:00:00 2001 From: jinyx5 Date: Tue, 1 Sep 2026 00:45:14 +0800 Subject: [PATCH 2/7] fix(distributed): default device mesh to the detected platform create_device_mesh_from_config defaulted device_type to "cuda" and 12 of 14 call sites rely on the default, so parallel denoising on NPU crashed inside torch DeviceMesh with "integer modulo by zero" (zero visible CUDA devices). Resolve the default from current_platform; explicit arguments keep working. Verified: 4-card (cfg=2 x ulysses=2) Wan2.2-TI2V-5B run on Ascend 910B2 builds the [cfg, ulysses] mesh over hccl; on CUDA the default resolves to "cuda" as before. --- telefuser/distributed/device_mesh.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/telefuser/distributed/device_mesh.py b/telefuser/distributed/device_mesh.py index 8281327e..ce3e69ff 100644 --- a/telefuser/distributed/device_mesh.py +++ b/telefuser/distributed/device_mesh.py @@ -13,10 +13,11 @@ from torch.distributed.device_mesh import DeviceMesh from telefuser.core.config import ParallelConfig +from telefuser.platforms import current_platform from telefuser.utils.logging import logger -def create_device_mesh_from_config(parallel_config: ParallelConfig, device_type: str = "cuda") -> DeviceMesh: +def create_device_mesh_from_config(parallel_config: ParallelConfig, device_type: str | None = None) -> DeviceMesh: """Create PyTorch DeviceMesh from ParallelConfig. Mesh dimensions are built in order: DP -> CFG -> SP (ring, ulysses) -> PP -> TP @@ -24,11 +25,13 @@ def create_device_mesh_from_config(parallel_config: ParallelConfig, device_type: Args: parallel_config: Parallel configuration with degrees for each dimension - device_type: Device type ("cuda" or "cpu") + device_type: Device type ("cuda", "npu", or "cpu"); defaults to the current platform's device type Returns: PyTorch DeviceMesh instance with named dimensions """ + if device_type is None: + device_type = current_platform.device_type _validate_parallel_config(parallel_config) sp_degree = parallel_config.sp_ulysses_degree * parallel_config.sp_ring_degree From f7d693d34dff82986a57c05c4d7abaceeb907c65 Mon Sep 17 00:00:00 2001 From: jinyx5 Date: Tue, 1 Sep 2026 00:45:14 +0800 Subject: [PATCH 3/7] fix(worker): marshal parallel-worker queue tensors through CPU off CUDA Parallel workers exchange request and result tensors over torch.multiprocessing queues, which relies on device IPC. torch_npu cross-process sharing raised "devptr INTERNAL ASSERT FAILED ... entry in cache has missing shared_ptr" when rebuilding queued NPU tensors. Enable the existing queue_with_cpu marshalling by default on non-CUDA platforms, mirror it on the result path, and move results back to the stage device in the main process. Worker dispatch already moves inputs to the local device, so the contract is unchanged and CUDA keeps direct device-queue transport. Verified: 4-card Wan2.2-TI2V-5B denoising on Ascend 910B2 completes with workers exiting cleanly. --- telefuser/worker/parallel_worker.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/telefuser/worker/parallel_worker.py b/telefuser/worker/parallel_worker.py index 2c160544..9adb8450 100644 --- a/telefuser/worker/parallel_worker.py +++ b/telefuser/worker/parallel_worker.py @@ -158,6 +158,9 @@ def _worker_loop( y = tensor_output_channel.send(y) # Always output results when world_size=1 if world_size == 1 or rank == 0: + if current_platform.device_type != "cuda": + # Queue transport without reliable device IPC: marshal results through CPU. + y = to_device(y, "cpu") queue_out.put(y) except Exception as e: import traceback @@ -206,7 +209,8 @@ def __init__( self.device_ids = list(range(self.world_size)) self.name: str = f"Parallel Worker {stage.name}" - self.queue_with_cpu: bool = parallel_config.queue_with_cpu + # Queue transport requires CPU marshalling on platforms without reliable device IPC (e.g. NPU). + self.queue_with_cpu: bool = parallel_config.queue_with_cpu or current_platform.device_type != "cuda" self.timeout: int = parallel_config.timeout self._lifecycle_lock = threading.Lock() self._failed = False @@ -319,6 +323,9 @@ def _wait_result(self, method_name: str) -> Any: reason = f"{method_name} failed: {result}" self._mark_failed(reason) raise RuntimeError(f"ParallelWorker:{self.name} {reason}") from result + if current_platform.device_type != "cuda": + # CPU-marshalled queue results move back to the stage device for downstream consumers. + result = to_device(result, self._stage.device) return result def enable_metrics(self, registry: Any | None = None) -> None: From 79ad1762d5430164a63cb6ffc6679b8e043c8f67 Mon Sep 17 00:00:00 2001 From: jinyx5 Date: Tue, 1 Sep 2026 00:45:22 +0800 Subject: [PATCH 4/7] refactor(ops): move fused FP8 QKV Triton kernels into kernel.triton ops/fp8_attention.py defined @triton.jit kernels at module scope, making "import triton" unconditional for every consumer of the wan_video and minimax DiT import chains even though the fused path is CUDA-only by contract. Move the two kernels and their launcher to telefuser/kernel/triton/fp8_attention.py and import them lazily inside quantize_fp8_qkv after its existing CUDA validation, matching the established ops -> kernel.triton dispatch pattern (see ops/rotary.py and kernel/__init__.py). The pure-torch quantize/dequantize helpers are unchanged. Verified: wan22 pipeline import and Wan2.2-TI2V-5B smoke succeed on an NPU host without triton installed; tests/unit/models/test_wan_video_sol_attention.py passes with CUDA-only cases skipped. --- telefuser/kernel/triton/fp8_attention.py | 119 +++++++++++++++++++++++ telefuser/ops/fp8_attention.py | 110 +-------------------- 2 files changed, 123 insertions(+), 106 deletions(-) create mode 100644 telefuser/kernel/triton/fp8_attention.py diff --git a/telefuser/kernel/triton/fp8_attention.py b/telefuser/kernel/triton/fp8_attention.py new file mode 100644 index 00000000..a620f374 --- /dev/null +++ b/telefuser/kernel/triton/fp8_attention.py @@ -0,0 +1,119 @@ +"""Fused block-scaled FP8 Q/K/V quantization Triton kernels.""" + +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _quantize_qkv_fp8_stage1( + q, + k, + v, + q_out, + k_out, + q_scale, + k_scale, + v_scale, + tokens: tl.constexpr, + heads: tl.constexpr, + head_dim: tl.constexpr, + block: tl.constexpr, +): + block_idx = tl.program_id(0) + batch_head = tl.program_id(1) + batch = batch_head // heads + head = batch_head % heads + token_offsets = block_idx * block + tl.arange(0, block) + dim_offsets = tl.arange(0, head_dim) + valid = token_offsets < tokens + offsets = ((batch * tokens + token_offsets[:, None]) * heads + head) * head_dim + dim_offsets[None, :] + q_values = tl.load(q + offsets, mask=valid[:, None], other=0.0).to(tl.float32) + k_values = tl.load(k + offsets, mask=valid[:, None], other=0.0).to(tl.float32) + v_values = tl.load(v + offsets, mask=valid[:, None], other=0.0).to(tl.float32) + + q_s = tl.maximum(tl.max(tl.max(tl.abs(q_values), axis=1), axis=0), 1.0e-6) / 448.0 + k_s = tl.maximum(tl.max(tl.max(tl.abs(k_values), axis=1), axis=0), 1.0e-6) / 448.0 + scale_offset = (batch * tl.cdiv(tokens, block) + block_idx) * heads + head + tl.store(q_scale + scale_offset, q_s) + tl.store(k_scale + scale_offset, k_s) + tl.store(q_out + offsets, q_values / q_s, mask=valid[:, None]) + tl.store(k_out + offsets, k_values / k_s, mask=valid[:, None]) + + v_s = tl.max(tl.abs(v_values), axis=0) / 448.0 + v_scale_offsets = (batch * heads + head) * head_dim + dim_offsets + tl.atomic_max(v_scale + v_scale_offsets, v_s) + + +@triton.jit +def _quantize_qkv_fp8_stage2_v( + v, + v_out, + v_scale, + tokens: tl.constexpr, + heads: tl.constexpr, + head_dim: tl.constexpr, + block: tl.constexpr, +): + block_idx = tl.program_id(0) + batch_head = tl.program_id(1) + batch = batch_head // heads + head = batch_head % heads + token_offsets = block_idx * block + tl.arange(0, block) + dim_offsets = tl.arange(0, head_dim) + valid = token_offsets < tokens + input_offsets = ((batch * tokens + token_offsets[:, None]) * heads + head) * head_dim + dim_offsets[None, :] + output_offsets = ((batch * heads + head) * head_dim + dim_offsets[None, :]) * tokens + token_offsets[:, None] + scale_offsets = (batch * heads + head) * head_dim + dim_offsets + scale = tl.maximum(tl.load(v_scale + scale_offsets), 1.0e-6 / 448.0) + values = tl.load(v + input_offsets, mask=valid[:, None], other=0.0).to(tl.float32) + tl.store(v_out + output_offsets, values / scale[None, :], mask=valid[:, None]) + + +def quantize_fp8_qkv_triton( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + block_size: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Launch the fused Q/K/V quantization kernels on validated CUDA inputs.""" + batch, tokens, heads, head_dim = q.shape + blocks = triton.cdiv(tokens, block_size) + q_out = torch.empty(q.shape, device=q.device, dtype=torch.float8_e4m3fn) + k_out = torch.empty_like(q_out) + v_storage = torch.empty((batch, heads, head_dim, tokens), device=q.device, dtype=torch.float8_e4m3fn) + q_scale = torch.empty((batch, blocks, heads), device=q.device, dtype=torch.float32) + k_scale = torch.ones_like(q_scale) + v_scale = torch.zeros((batch, heads, head_dim), device=q.device, dtype=torch.float32) + grid = (blocks, batch * heads) + _quantize_qkv_fp8_stage1[grid]( + q, + k, + v, + q_out, + k_out, + q_scale, + k_scale, + v_scale, + tokens, + heads, + head_dim, + block_size, + num_warps=8, + num_stages=1, + ) + _quantize_qkv_fp8_stage2_v[grid]( + v, + v_storage, + v_scale, + tokens, + heads, + head_dim, + block_size, + num_warps=8, + num_stages=1, + ) + v_out = v_storage.permute(0, 3, 1, 2) + return q_out, k_out, v_out, q_scale, k_scale, v_scale diff --git a/telefuser/ops/fp8_attention.py b/telefuser/ops/fp8_attention.py index 7a2f36d5..eeec0f20 100644 --- a/telefuser/ops/fp8_attention.py +++ b/telefuser/ops/fp8_attention.py @@ -4,77 +4,10 @@ import torch import torch.nn.functional as F -import triton -import triton.language as tl FP8_ATTENTION_BLOCK_SIZE = 64 -@triton.jit -def _quantize_qkv_fp8_stage1( - q, - k, - v, - q_out, - k_out, - q_scale, - k_scale, - v_scale, - tokens: tl.constexpr, - heads: tl.constexpr, - head_dim: tl.constexpr, - block: tl.constexpr, -): - block_idx = tl.program_id(0) - batch_head = tl.program_id(1) - batch = batch_head // heads - head = batch_head % heads - token_offsets = block_idx * block + tl.arange(0, block) - dim_offsets = tl.arange(0, head_dim) - valid = token_offsets < tokens - offsets = ((batch * tokens + token_offsets[:, None]) * heads + head) * head_dim + dim_offsets[None, :] - q_values = tl.load(q + offsets, mask=valid[:, None], other=0.0).to(tl.float32) - k_values = tl.load(k + offsets, mask=valid[:, None], other=0.0).to(tl.float32) - v_values = tl.load(v + offsets, mask=valid[:, None], other=0.0).to(tl.float32) - - q_s = tl.maximum(tl.max(tl.max(tl.abs(q_values), axis=1), axis=0), 1.0e-6) / 448.0 - k_s = tl.maximum(tl.max(tl.max(tl.abs(k_values), axis=1), axis=0), 1.0e-6) / 448.0 - scale_offset = (batch * tl.cdiv(tokens, block) + block_idx) * heads + head - tl.store(q_scale + scale_offset, q_s) - tl.store(k_scale + scale_offset, k_s) - tl.store(q_out + offsets, q_values / q_s, mask=valid[:, None]) - tl.store(k_out + offsets, k_values / k_s, mask=valid[:, None]) - - v_s = tl.max(tl.abs(v_values), axis=0) / 448.0 - v_scale_offsets = (batch * heads + head) * head_dim + dim_offsets - tl.atomic_max(v_scale + v_scale_offsets, v_s) - - -@triton.jit -def _quantize_qkv_fp8_stage2_v( - v, - v_out, - v_scale, - tokens: tl.constexpr, - heads: tl.constexpr, - head_dim: tl.constexpr, - block: tl.constexpr, -): - block_idx = tl.program_id(0) - batch_head = tl.program_id(1) - batch = batch_head // heads - head = batch_head % heads - token_offsets = block_idx * block + tl.arange(0, block) - dim_offsets = tl.arange(0, head_dim) - valid = token_offsets < tokens - input_offsets = ((batch * tokens + token_offsets[:, None]) * heads + head) * head_dim + dim_offsets[None, :] - output_offsets = ((batch * heads + head) * head_dim + dim_offsets[None, :]) * tokens + token_offsets[:, None] - scale_offsets = (batch * heads + head) * head_dim + dim_offsets - scale = tl.maximum(tl.load(v_scale + scale_offsets), 1.0e-6 / 448.0) - values = tl.load(v + input_offsets, mask=valid[:, None], other=0.0).to(tl.float32) - tl.store(v_out + output_offsets, values / scale[None, :], mask=valid[:, None]) - - def quantize_fp8_qkv( q: torch.Tensor, k: torch.Tensor, @@ -86,46 +19,11 @@ def quantize_fp8_qkv( raise ValueError("q, k, and v must share shape [B, T, H, D]") if not (q.is_cuda and q.is_contiguous() and k.is_contiguous() and v.is_contiguous()): raise ValueError("fused FP8 QKV quantization requires contiguous CUDA tensors") - batch, tokens, heads, head_dim = q.shape - if head_dim != 128: + if q.shape[-1] != 128: raise ValueError("fused FP8 QKV quantization requires head dimension 128") - blocks = triton.cdiv(tokens, FP8_ATTENTION_BLOCK_SIZE) - q_out = torch.empty(q.shape, device=q.device, dtype=torch.float8_e4m3fn) - k_out = torch.empty_like(q_out) - v_storage = torch.empty((batch, heads, head_dim, tokens), device=q.device, dtype=torch.float8_e4m3fn) - q_scale = torch.empty((batch, blocks, heads), device=q.device, dtype=torch.float32) - k_scale = torch.ones_like(q_scale) - v_scale = torch.zeros((batch, heads, head_dim), device=q.device, dtype=torch.float32) - grid = (blocks, batch * heads) - _quantize_qkv_fp8_stage1[grid]( - q, - k, - v, - q_out, - k_out, - q_scale, - k_scale, - v_scale, - tokens, - heads, - head_dim, - FP8_ATTENTION_BLOCK_SIZE, - num_warps=8, - num_stages=1, - ) - _quantize_qkv_fp8_stage2_v[grid]( - v, - v_storage, - v_scale, - tokens, - heads, - head_dim, - FP8_ATTENTION_BLOCK_SIZE, - num_warps=8, - num_stages=1, - ) - v_out = v_storage.permute(0, 3, 1, 2) - return q_out, k_out, v_out, q_scale, k_scale, v_scale + from telefuser.kernel.triton.fp8_attention import quantize_fp8_qkv_triton + + return quantize_fp8_qkv_triton(q, k, v, FP8_ATTENTION_BLOCK_SIZE) def quantize_fp8_per_block( From 162b9e344260220ea15901cdcfee4c571ed8fdc9 Mon Sep 17 00:00:00 2001 From: jinyx5 Date: Tue, 1 Sep 2026 00:45:22 +0800 Subject: [PATCH 5/7] fix(distributed): allocate pipeline P2P buffers on the platform device recv, recv_latent, and recv_latent_async allocated receive buffers with device="cuda", which fails on NPU-only hosts before irecv can run. Allocate on current_platform.device_type (identical behavior on CUDA) and parameterize the test tensor device the same way so the suite runs on CUDA, NPU, and CPU hosts. Verified: tests/unit/distributed/test_pp_comm.py 12 passed on Ascend 910B2 (previously 6 device-related failures). --- telefuser/distributed/pp_comm.py | 7 ++++--- tests/unit/distributed/test_pp_comm.py | 18 +++++++++++------- 2 files changed, 15 insertions(+), 10 deletions(-) diff --git a/telefuser/distributed/pp_comm.py b/telefuser/distributed/pp_comm.py index 9c3b9085..5a4a2de7 100644 --- a/telefuser/distributed/pp_comm.py +++ b/telefuser/distributed/pp_comm.py @@ -24,6 +24,7 @@ import torch import torch.distributed as dist +from telefuser.platforms import current_platform from telefuser.utils.logging import logger @@ -107,7 +108,7 @@ def recv( if buffer is None: if shape is None: raise ValueError("Either buffer or shape must be provided") - buffer = torch.empty(shape, dtype=torch.float16, device="cuda") + buffer = torch.empty(shape, dtype=torch.float16, device=current_platform.device_type) buffer = buffer.contiguous() if async_op: @@ -266,7 +267,7 @@ def recv_latent(self, shape: tuple | None = None, dtype: torch.dtype = torch.bfl if shape is None: raise ValueError("recv_latent: shape must be provided") - buffer = torch.empty(shape, dtype=dtype, device="cuda") + buffer = torch.empty(shape, dtype=dtype, device=current_platform.device_type) buffer = buffer.contiguous() work = dist.irecv(buffer, self.recv_src, group=self._process_group) work.wait() @@ -300,7 +301,7 @@ def recv_latent_async(self, shape: tuple, dtype: torch.dtype = torch.bfloat16) - if self.is_first_stage: raise RuntimeError("recv_latent_async: First stage has no previous stage to receive from") - buffer = torch.empty(shape, dtype=dtype, device="cuda") + buffer = torch.empty(shape, dtype=dtype, device=current_platform.device_type) buffer = buffer.contiguous() work = dist.irecv(buffer, self.recv_src, group=self._process_group) return buffer, work diff --git a/tests/unit/distributed/test_pp_comm.py b/tests/unit/distributed/test_pp_comm.py index f247d001..99b52a7d 100644 --- a/tests/unit/distributed/test_pp_comm.py +++ b/tests/unit/distributed/test_pp_comm.py @@ -7,10 +7,14 @@ import torch.distributed as dist from telefuser.distributed.pp_comm import PipelineP2PComm +from telefuser.platforms import current_platform # Skip if distributed not available HAS_DISTRIBUTED = dist.is_available() +# Buffers follow the detected platform so the suite runs on CUDA, NPU, and CPU hosts alike. +DEVICE = current_platform.device_type + pytestmark = [ pytest.mark.skipif(not HAS_DISTRIBUTED, reason="Distributed not available"), pytest.mark.distributed, @@ -56,7 +60,7 @@ def test_send_on_last_stage_logs_warning(self): """Test that send on last stage logs warning and returns None.""" comm = PipelineP2PComm(None) # Single GPU, is_last_stage=True - tensor = torch.randn(1, 10, 512, device="cuda") + tensor = torch.randn(1, 10, 512, device=DEVICE) result = comm.send(tensor) assert result is None @@ -72,7 +76,7 @@ def test_send_recv_single_gpu(self): """Test send_recv on single GPU (no-op).""" comm = PipelineP2PComm(None) # Single GPU - send_tensor = torch.randn(1, 10, 512, device="cuda") + send_tensor = torch.randn(1, 10, 512, device=DEVICE) result = comm.send_recv(send_tensor) # On single GPU, recv_buffer is None since there's no previous stage @@ -114,7 +118,7 @@ def test_queue_send_on_last_stage(self): """Test queue_send on last stage does nothing.""" comm = PipelineP2PComm(None) # is_last_stage=True - tensor = torch.randn(1, 10, 512, device="cuda") + tensor = torch.randn(1, 10, 512, device=DEVICE) comm.queue_send(tensor) assert len(comm._ops) == 0 @@ -123,7 +127,7 @@ def test_queue_recv_on_first_stage(self): """Test queue_recv on first stage does nothing.""" comm = PipelineP2PComm(None) # is_first_stage=True - buffer = torch.randn(1, 10, 512, device="cuda") + buffer = torch.randn(1, 10, 512, device=DEVICE) comm.queue_recv(buffer) assert len(comm._ops) == 0 @@ -152,7 +156,7 @@ def test_send_latent_on_last_stage(self): """Test send_latent on last stage returns early.""" comm = PipelineP2PComm(None) # is_last_stage=True - tensor = torch.randn(1, 10, 512, device="cuda") + tensor = torch.randn(1, 10, 512, device=DEVICE) # Should not raise, just return comm.send_latent(tensor) @@ -182,7 +186,7 @@ def test_send_latent_async_on_last_stage(self): """Test send_latent_async on last stage returns None.""" comm = PipelineP2PComm(None) # is_last_stage=True - tensor = torch.randn(1, 10, 512, device="cuda") + tensor = torch.randn(1, 10, 512, device=DEVICE) result = comm.send_latent_async(tensor) assert result is None @@ -248,7 +252,7 @@ def test_actual_p2p_communication(self): if comm.is_first_stage: # Send tensor to next stage - send_tensor = torch.ones(1, 10, 512, device="cuda") * comm.rank + send_tensor = torch.ones(1, 10, 512, device=DEVICE) * comm.rank comm.send_latent(send_tensor) elif comm.is_last_stage: # Receive tensor from previous stage From 20e1038967cf4647cd2a7a15bc04bf2d260fbe01 Mon Sep 17 00:00:00 2001 From: jinyx5 Date: Tue, 1 Sep 2026 09:11:53 +0800 Subject: [PATCH 6/7] fix(worker): keep marshalled queue results device-agnostic The non-CUDA result path introduced for CPU-marshalled queues moved results back to self._stage.device inside _wait_result. That breaks worker unit tests on non-CUDA hosts (mocked stages resolve .device through the method proxy or to a MagicMock), and it is unnecessary: stage entry points already place their inputs, and the Wan VAE decode paths move latents themselves. Drop the move so _wait_result returns results untouched; workers still marshal results through CPU when device IPC is unavailable. Verified: tests/unit/worker passes on an Ascend host (same non-CUDA branch as CPU CI) and the 4-card Wan2.2-TI2V-5B smoke still completes; ruff clean. --- telefuser/worker/parallel_worker.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/telefuser/worker/parallel_worker.py b/telefuser/worker/parallel_worker.py index 9adb8450..01727299 100644 --- a/telefuser/worker/parallel_worker.py +++ b/telefuser/worker/parallel_worker.py @@ -323,9 +323,6 @@ def _wait_result(self, method_name: str) -> Any: reason = f"{method_name} failed: {result}" self._mark_failed(reason) raise RuntimeError(f"ParallelWorker:{self.name} {reason}") from result - if current_platform.device_type != "cuda": - # CPU-marshalled queue results move back to the stage device for downstream consumers. - result = to_device(result, self._stage.device) return result def enable_metrics(self, registry: Any | None = None) -> None: From 8a4220e3b3e92d0fd18a939566d69fb179f1eae0 Mon Sep 17 00:00:00 2001 From: jinyx5 Date: Tue, 1 Sep 2026 11:30:50 +0800 Subject: [PATCH 7/7] feat(examples): auto-detect device in the wan22 5B T2V example The example hardcoded device="cuda", so running it on an NPU host required editing the script. Default the pipeline device to current_platform.device_type so the same command runs unmodified on CUDA and NPU; on CUDA hosts this resolves to "cuda" as before. Verified: python examples/wan_video/wan22_t2v_5b.py --gpu_num 1 and --gpu_num 4 run unmodified on Ascend 910B (50 steps, 121 frames, default prompt). --- examples/wan_video/wan22_t2v_5b.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/examples/wan_video/wan22_t2v_5b.py b/examples/wan_video/wan22_t2v_5b.py index 76cc1a56..573f1eb5 100644 --- a/examples/wan_video/wan22_t2v_5b.py +++ b/examples/wan_video/wan22_t2v_5b.py @@ -17,6 +17,7 @@ Wan22TI2VPipeline, Wan22TI2VPipelineConfig, ) +from telefuser.platforms import current_platform from telefuser.utils.utils import get_example_name from telefuser.utils.video import get_target_video_size_from_ratio, save_video @@ -80,7 +81,7 @@ def get_pipeline(parallelism: int = 1, model_root: str = PPL_CONFIG["model_root" ) # Create pipeline - pipe = Wan22TI2VPipeline(device="cuda", torch_dtype=torch.bfloat16) + pipe = Wan22TI2VPipeline(device=current_platform.device_type, torch_dtype=torch.bfloat16) # Configure pipeline pipe_config = Wan22TI2VPipelineConfig()