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() 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 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/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/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/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) 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( diff --git a/telefuser/worker/parallel_worker.py b/telefuser/worker/parallel_worker.py index 2c160544..01727299 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 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