已关闭
[Bug]: 自定义算子的python callback特性不能正确处理可选参数为None和参数类型为torch.dtype的输入 #799
Chang-an-HW创建于  13 天前关闭于  11 天前
Chang-an-HW成员
13 天前 创建

在提交问题之前,请通过搜索现有和历史问题确保该问题尚未被提出并解决。

您的环境信息

-- CANN 版本 (e.g., 5.x.x, 8.x.x):  
-- Pytorch/Torch_npu 版本 (e.g., v2.6.0, v2.9.0):
-- Python 版本 (e.g., Python3.9.x, Python3.11.x):
-- 操作系统版本 (e.g., Ubuntu 18.04, EulerOs 2.0):

🐛 请描述bug

#!/usr/bin/python3
# -*- coding: utf-8 -*-

import torch
import torch.nn as nn
import os
import triton
import triton.language as tl
import torch_npu
import torchair as tng
from torchair.configs.compiler_config import CompilerConfig

@triton.jit
def _scalar_mul_kernel_051(x_ptr, output_ptr, scalar_ptr, n_elements, BLOCK_SIZE: tl.constexpr):
    pid = tl.program_id(axis=0)
    block_start = pid * BLOCK_SIZE
    offsets = block_start + tl.arange(0, BLOCK_SIZE)
    mask = offsets < n_elements
    x = tl.load(x_ptr + offsets, mask=mask)
    scalar_val = tl.load(scalar_ptr)
    output = x * scalar_val
    tl.store(output_ptr + offsets, output, mask=mask)

@torch.library.triton_op("triton_hop::scalar_mul_unsupported_051", mutates_args={})
def _triton_op_scalar_mul_051(x: torch.Tensor, scalar_tensor: torch.Tensor) -> torch.Tensor:
    output = torch.empty_like(x)
    n_elements = output.numel()
    grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
    torch.library.wrap_triton(_scalar_mul_kernel_051)[grid](x, output, scalar_tensor, n_elements, BLOCK_SIZE=1024)
    return output


class _UnsupportedPlacementModel051(nn.Module):
    def __init__(self, dim=64):
        super().__init__()
        self.dim = dim
        self.gamma = nn.Parameter(torch.ones(dim))
        self.beta = nn.Parameter(torch.zeros(dim))
        self.scalar = nn.Parameter(torch.ones(1))
        self.smooth_scales = nn.Parameter(torch.ones(1, dtype=torch.float32))
        self.offsets = nn.Parameter(torch.zeros(1, dtype=torch.float32))

    def forward(self, x):
        def chain_out(o):
            if isinstance(o, (tuple, list)):
                o = o[0]
            f = o.float().reshape(-1)
            if f.numel() >= 4096:
                return f[:4096].reshape(64, 64)
            return torch.cat([f, torch.zeros(4096 - f.numel(), device=f.device)]).reshape(64, 64)

        h = x.view(64, 64)
        h = _triton_op_scalar_mul_051(h, self.scalar)
        h, _, _ = torch_npu.npu_group_norm_silu(h, self.gamma, self.beta, group=4)
        h = torch_npu.npu_gelu(h)
        h = torch.nn.functional.relu(h)
        h = h.pow(2)
        h = h + 1.0
        h = h * self.gamma
        h = h - 0.5
        h = h.abs()
        h = h.sqrt()
        h = h / 4.0

        h = chain_out(torch_npu.npu_gelu_mul(h, approximate="none"))
        h = chain_out(torch_npu.npu_clipped_swiglu(h, alpha=1.702, limit=7.0, bias=1.0, interleaved=True))
        h = chain_out(torch_npu.npu_swiglu_quant(h.half(), smooth_scales=self.smooth_scales, offsets=self.offsets, activate_left=False, quant_mode=0))

        return h


model = _UnsupportedPlacementModel051().npu()
x_npu = torch.randn(4096, dtype=torch.float32).npu()

config = CompilerConfig()
npu_backend = tng.get_npu_backend(compiler_config=config)
compiled = torch.compile(model, fullgraph=True, backend=npu_backend, dynamic=True)
with torch.no_grad():
    graph_out = compiled(x_npu)

npu_swiglu_quant算子有一个名为dst_type默认值为None,类型为torch.dtype的可选输入参数,执行这个用例走到生成的Converter时需要将dst_type输入转为ge int型Const,但是python无法将None和torch.dtype值转换为int,程序崩溃

likedislike
CChang-an-HW成员
13 天前 添加了label:bug
CChang-an-HW成员
13 天前 添加了label:bug
CChang-an-HW成员
13 天前 关联了pull request:增强可选输入处理,确保None值正确映射为缺失输入,添加dtype值归一化逻辑
wj1ewj1e成员
13 天前 将 Chang-an-HW 设为负责人
ascend-robotascend-robot成员
11 天前 关闭了 issue
ascend-robotascend-robot成员
11 天前 issue状态由 TODO 改变为 DONE
ascend-robotascend-robot成员
11 天前 添加了label:resolved