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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions docs/en/installation.md
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ distribution are installed separately.
| PyTorch | 2.6 or newer |
| CUDA toolkit | 12.8 or newer for the maintained CUDA development path |
| ROCm | 7.x with a PyTorch `+rocm` build for AMD GPUs; see the ROCm note under verification |
| Ascend NPU | CANN 8.2 with matching `torch` and `torch_npu` builds; see the NPU note under verification |
| GPU | Depends on the selected model; check its Cookbook guide |

An example may impose stricter versions or GPU architecture requirements. In particular, locally built `tf-kernel`
Expand Down Expand Up @@ -62,6 +63,11 @@ ROCm before CUDA. Examples ending in `_rocm.py` (for example `examples/wan_video
are the validated entry points; see [Hardware Platforms](platforms.md) for per-platform capabilities and backend
availability.

On Huawei Ascend hosts, install a `torch_npu` build matching your PyTorch version instead of the CUDA toolkit
path. The platform layer selects NPU when `torch_npu` imports and an Ascend device is visible
(`ASCEND_RT_VISIBLE_DEVICES` controls visibility); `torch.cuda.is_available()` printing `False` is expected there.
`examples/wan_video/wan22_t2v_5b.py` is the validated NPU entry point.

## Model Checkpoints

TeleFuser does not bundle model weights. The [Supported Models](supported_models.md) page links to each Cookbook
Expand Down
23 changes: 19 additions & 4 deletions docs/en/platforms.md
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,22 @@ setup path.

## NPU and CPU

The NPU platform targets Huawei Ascend devices through `torch_npu` with the HCCL distributed backend. It is wired
into the platform and ops dispatch layers, but the maintained examples are validated on CUDA and, for select
examples, ROCm — validate on your target NPU before production use. The CPU platform is the fallback when no
accelerator is detected; it is intended for tests and for pipelines that explicitly request CPU execution.
The NPU platform targets Huawei Ascend devices through `torch_npu` with the HCCL distributed backend.

- Attention uses `TORCH_SDPA` through the native fallback paths; no `tf-kernel`, `flash_attn`, `sageattention`, or
`triton` installation is required.
- The ops layer selects `forward_npu` where an op defines one and otherwise falls back to native PyTorch. No
NPU-optimized kernels are integrated yet, so pipelines run entirely on the native paths.
- Multi-card inference uses HCCL through the `hccl` backend string. Parallel-worker queues marshal tensors through
CPU because `torch_npu` has no reliable cross-process device IPC, and each spawned worker group receives a
distinct `HCCL_IF_BASE_PORT` so concurrent groups do not collide on HCCL's data-plane socket range.
- `torch.compile` is not validated on NPU; NPU examples run eager.
- The validated entry point is
[Wan2.2 TI2V-5B text-to-video](https://github.com/Tele-AI/TeleFuser/tree/main/examples/wan_video)
(`wan22_t2v_5b.py`), which auto-detects the platform and runs unmodified on an Atlas 910B (CANN 8.2, torch 2.9
with a matching `torch_npu`): single-card, and four-card CFG × Ulysses parallelism over HCCL. Wan2.2 A14B shares
these code paths but has not been exercised on NPU hardware; validate other examples on your target NPU before
production use.

The CPU platform is the fallback when no accelerator is detected; it is intended for tests and for pipelines that
explicitly request CPU execution.
5 changes: 5 additions & 0 deletions docs/zh/installation.md
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
| PyTorch | 2.6 或更高版本 |
| CUDA Toolkit | 当前 CUDA 开发路径要求 12.8 或更高版本 |
| ROCm | AMD GPU 使用 ROCm 7.x 与 PyTorch `+rocm` 构建,详见验证安装一节的说明 |
| 昇腾 NPU | CANN 8.2 搭配版本匹配的 `torch` 与 `torch_npu` 构建,详见验证安装一节的说明 |
| GPU | 取决于所选模型,以对应 Cookbook 为准 |

具体示例可能要求更严格的软件版本或 GPU 架构。特别是本地构建的 `tf-kernel` 产物与其记录的 PyTorch、
Expand Down Expand Up @@ -61,6 +62,10 @@ AMD ROCm 主机应安装 PyTorch `+rocm` 构建,而非 CUDA Toolkit 路径。H
`examples/wan_video/wan21_1_3b_text_to_video_rocm.py`)是已验证的入口;各平台能力与后端可用性见
[硬件平台](platforms.md)。

华为昇腾主机应安装与 PyTorch 版本匹配的 `torch_npu`,而非 CUDA Toolkit 路径。当 `torch_npu` 可导入且存在
可见昇腾设备(由 `ASCEND_RT_VISIBLE_DEVICES` 控制)时,平台层会选择 NPU;此时 `torch.cuda.is_available()`
输出 `False` 属预期行为。`examples/wan_video/wan22_t2v_5b.py` 是已验证的 NPU 入口。

## 模型权重

TeleFuser 不随软件包分发模型权重。[支持的模型](supported_models.md)页面会链接到各模型的 Cookbook,
Expand Down
19 changes: 16 additions & 3 deletions docs/zh/platforms.md
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,19 @@ ROCm 支持面向使用 ROCm 7.x 与 PyTorch `+rocm` 构建的 AMD GPU;安装

## NPU 与 CPU

NPU 平台通过 `torch_npu` 与 HCCL 分布式后端支持华为昇腾设备。平台层与算子分发层均已接入 NPU,但现有
示例在 CUDA 上验证、部分示例在 ROCm 上验证 —— 生产使用前请先在目标 NPU 上完成验证。CPU 平台是未检测到
加速器时的回退,面向测试以及显式请求 CPU 执行的 Pipeline。
NPU 平台通过 `torch_npu` 与 HCCL 分布式后端支持华为昇腾设备。

- 注意力经原生回退路径使用 `TORCH_SDPA`;无需安装 `tf-kernel`、`flash_attn`、`sageattention` 或 `triton`。
- 算子层在算子定义了 `forward_npu` 时优先选择,否则回退到 PyTorch 原生实现。目前尚未集成 NPU 优化内核,
Pipeline 完全运行在原生路径上。
- 多卡推理通过 `hccl` 后端字符串使用 HCCL。并行 worker 队列经 CPU 中转张量(`torch_npu` 不提供可靠的
跨进程设备 IPC);每个新起的 worker 组会分配独立的 `HCCL_IF_BASE_PORT`,避免并发组在 HCCL 数据面端口段
上冲突。
- `torch.compile` 在 NPU 上未验证;NPU 示例以 eager 模式运行。
- 已验证入口为
[Wan2.2 TI2V-5B 文生视频](https://github.com/Tele-AI/TeleFuser/tree/main/examples/wan_video)
(`wan22_t2v_5b.py`):示例自动检测平台、免修改运行,已在 Atlas 910B(CANN 8.2,torch 2.9 搭配版本匹配的
`torch_npu`)上验证单卡以及 4 卡 CFG × Ulysses 并行(HCCL)。Wan2.2 A14B 复用同一代码路径,但尚未在
NPU 硬件上运行;其他示例在生产使用前请先在目标 NPU 上完成验证。

CPU 平台是未检测到加速器时的回退,面向测试以及显式请求 CPU 执行的 Pipeline。
6 changes: 6 additions & 0 deletions examples/wan_video/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,9 @@ Video generation using Wan2.1 and Wan2.2 models for Text-to-Video and Image-to-V
- GPU: AMD ROCm GPUs for scripts ending in `_rocm.py`; validated on a Radeon RX 9070 (ROCm 7.2, `torch` built with
`+rocm`). These examples use the PyTorch SDPA attention backend and need no tf-kernel, flash-attn, or SageAttention
installation
- GPU: Huawei Ascend NPUs for `wan22_t2v_5b.py`, which auto-detects the platform; validated on an Atlas 910B
(CANN 8.2, torch 2.9 with a matching `torch_npu`). NPU execution uses the PyTorch SDPA attention backend and
needs no tf-kernel, flash-attn, SageAttention, or triton installation
- Software: the standard TeleFuser installation; optional attention, FP8, Ray, and RIFE paths require their respective
dependencies
- Input assets: a readable image for I2V/FL2V and optional LoRA, distillation, cache, or RIFE weights for those variants
Expand Down Expand Up @@ -463,6 +466,9 @@ python examples/wan_video/wan22_t2v_5b.py --resolution 480p --aspect_ratio 16:9
- CFG parallel enabled by default (cfg_scale=5.0)
- Ulysses sequence parallelism for multi-GPU
- 50-step UNPC sampling with sigma_shift=5.0
- Platform auto-detection via `current_platform`: the script runs unmodified on CUDA and Ascend NPU hosts
- Validated on an Ascend Atlas 910B (CANN 8.2, torch 2.9 with a matching `torch_npu`): single-card and 4-card
CFG × Ulysses over HCCL, PyTorch SDPA attention, eager execution

#### `wan22_14b_text_to_video_h100.py`

Expand Down
3 changes: 2 additions & 1 deletion examples/wan_video/wan22_t2v_5b.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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()
Expand Down
7 changes: 5 additions & 2 deletions telefuser/distributed/device_mesh.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,22 +13,25 @@
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
For USP (Unified Sequence Parallelism), ring and ulysses form a 2D sub-mesh.

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
Expand Down
7 changes: 4 additions & 3 deletions telefuser/distributed/pp_comm.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
import torch
import torch.distributed as dist

from telefuser.platforms import current_platform
from telefuser.utils.logging import logger


Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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
6 changes: 5 additions & 1 deletion telefuser/distributed/vae_spatial.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
119 changes: 119 additions & 0 deletions telefuser/kernel/triton/fp8_attention.py
Original file line number Diff line number Diff line change
@@ -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
6 changes: 5 additions & 1 deletion telefuser/models/wan22_video_vae.py
Original file line number Diff line number Diff line change
Expand Up @@ -1403,7 +1403,11 @@ def decode(
if self.parallelism > 1 and dist.is_initialized():
# tiled=True → tile_dist, tiled=False → 2d_split
method = "tile_dist" if tiled else "2d_split"
hidden_states_tensor = torch.stack(hidden_states)
# The stage passes a batched [B, C, T, H, W] tensor; lists of per-video tensors are stacked.
if isinstance(hidden_states, torch.Tensor):
hidden_states_tensor = hidden_states
else:
hidden_states_tensor = torch.stack(hidden_states)
return self.decode_parallel(hidden_states_tensor, device, method=method)

# Single GPU processing
Expand Down
6 changes: 5 additions & 1 deletion telefuser/models/wan_video_vae.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)


Expand Down
Loading
Loading