已关闭
[Bug]: 图模式下引入额外stride和view copy操作 #241
WenquanYang创建于  2025年12月29日关闭于  1月12日
WenquanYang
2025年12月29日 创建

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

您的环境信息

-- CANN 版本 (e.g., CANN 8.5.T8.0.B060):
-- Pytorch/Torch_npu 版本 (e.g., v2.6.0, v2.6.0.post3):
-- Python 版本 (e.g., Python 3.11):
-- 操作系统版本 (e.g., 欧拉os):

🐛 请描述bug

对tensor进行allgather后再做矩阵乘法,会引入额外的stride view copy操作
问题背景:

@torch.compile(backend=npu_backend, dynamic=False, fullgraph=True)
def opp(x, w_buffer, w_partial, w1, w2):
    dist.all_gather_into_tensor(w_buffer, w_partial)
    tng.scope.npu_wait_tensor(w1, w_buffer)
    h1 = torch.matmul(x, w1)
    o = torch.matmul(h1, w2)
    return o

对一份共享tensor进行allgather,两个weight指向这个共享tensor的不同部分,再对共享tensor做矩阵乘法后,两个weight单独进行矩阵乘法
eager模式下:profiling只有 allgather+matmul+matmul, 图模式会引入额外的操作

image.png

图结构dump后也能看到torchair阶段加入了很多额外的操作

完整代码片段如下:可用 torchrun --nproc-per-node 4 main.py执行

import os
import torch
import torch.distributed as dist
import torch_npu
import torchair
import torchair as tng
from torchair.configs.compiler_config import CompilerConfig

torch.manual_seed(0)
torchair.patch_for_hcom()

local_rank = int(os.environ["LOCAL_RANK"])
global_rank = int(os.environ["RANK"])
torch.npu.set_device(f"npu:{local_rank}")

dist.init_process_group(backend="hccl")
world_size = dist.get_world_size()

def print_rank0(*msg):
    if int(os.environ["RANK"]) == 0:
        print(*msg)

config = CompilerConfig()
config.experimental_config.topology_sorting_strategy = "StableRDFS"
# config.mode = "reduce-overhead"
if local_rank == 0:
    config.debug.graph_dump.type = "pbtxt"

npu_backend = torchair.get_npu_backend(compiler_config=config)

@torch.compile(backend=npu_backend, dynamic=False, fullgraph=True)
def opp(x, w_buffer, w_partial, w1, w2):
    dist.all_gather_into_tensor(w_buffer, w_partial)
    tng.scope.npu_wait_tensor(w1, w_buffer)
    h1 = torch.matmul(x, w1)
    o = torch.matmul(h1, w2)
    return o


def main():
    x = torch.randn(2, 8, device="npu", dtype=torch.bfloat16)
    wqk_buffer = torch.arange(8*16*2, device="npu", dtype=torch.bfloat16)
    wqk_partial = torch.arange(8*16*2//world_size, device="npu", dtype=torch.bfloat16)

    wq = wqk_buffer[:8*16].view(8, 16)
    wk = wqk_buffer[8*16:].view(16, 8)

    if local_rank == 0:
        PROFILING_PATH = "op_test"

        EXPERIMENTAL_CONFIG = torch_npu.profiler._ExperimentalConfig(
            aic_metrics=torch_npu.profiler.AiCMetrics.PipeUtilization,
            profiler_level=torch_npu.profiler.ProfilerLevel.Level1,
            l2_cache=False,
        )

        with torch_npu.profiler.profile(
            activities=[
                torch_npu.profiler.ProfilerActivity.CPU,
                torch_npu.profiler.ProfilerActivity.NPU,
            ],
            on_trace_ready=torch_npu.profiler.tensorboard_trace_handler(PROFILING_PATH),
            record_shapes=True,
            with_flops=False,
            with_modules=False,
            experimental_config=EXPERIMENTAL_CONFIG,
        ) as prof:
            o_bar = opp(x, wqk_buffer, wqk_partial, wq, wk)

    else:
        o_bar = opp(x, wqk_buffer, wqk_partial, wq, wk)

    torch.npu.synchronize()
    print_rank0(o_bar)

main()

likedislike
SunYapingSunYaping成员
2025年12月29日 将 tangjie66 设为负责人
tangjie66
tangjie66成员
2025年12月29日 评论:

您好,可以尝试将
wq = wqk_buffer[:816].view(8, 16)
wk = wqk_buffer[8
16:].view(16, 8)
改为
wq = wqk_buffer[:816].view(8, 16).clone()
wk = wqk_buffer[8
16:].view(16, 8).clone()
再看下图中有没有多余的viewcopy和strided等操作

以下是对应的解释:

  1. 为什么原来的写法有asstrided:
    以wq为例,原来的wq只是给wqk_buffer的前128个元素贴了个新标签,还是共享的一块内存,它的步长需要根据wqk_buffer的步长进行推导
    而npu需要在编译阶段,通过asstrided将布局规则写死,在执行时才能知道从wqk_buffer中怎么取数据。

  2. 为什么原来的写法中有viewcopy:
    all_gather_into_tensor执行的时候会往wqk_buffer中写新数据,wq和wqk_buffer共享一块内存
    想让wq能正确读取更新的数据,必须将输出的通信结果的内存布局通过viewcopy变成wq能识别的格式

  3. 为什么新写法没有这些算子:
    如果wq和wk不依赖其他张量,如wqk_buffer,它的内存是独立的,编译框架无需追踪两者的布局依赖

likedislike
WenquanYang
2025年12月30日 评论:

但是我这个场景里,需要wq, wk用wqk_buffer里的共享内存,有办法规避吗

likedislike
wang-pierre-jiacheng
wang-pierre-jiacheng成员
1月5日 评论:

但是我这个场景里,需要wq, wk用wqk_buffer里的共享内存,有办法规避吗

@WenquanYang
想保留这样的共享关系,可以考虑把w1和w2包含到fx图中。
否则torch.compile自行感知w1\w2与w_buffer的特殊关系,会导致图上多了as_strided、viewcopy

import os
import torch
import torch.distributed as dist
import torch_npu
import torchair
import torchair as tng
from torchair.configs.compiler_config import CompilerConfig

torch.manual_seed(0)
torchair.patch_for_hcom()

local_rank = int(os.environ["LOCAL_RANK"])
global_rank = int(os.environ["RANK"])
torch.npu.set_device(f"npu:{local_rank}")

dist.init_process_group(backend="hccl")
world_size = dist.get_world_size()

def print_rank0(*msg):
    if int(os.environ["RANK"]) == 0:
        print(*msg)

config = CompilerConfig()
config.experimental_config.topology_sorting_strategy = "StableRDFS"
# config.mode = "reduce-overhead"
if local_rank == 0:
    config.debug.graph_dump.type = "py"

npu_backend = torchair.get_npu_backend(compiler_config=config)

@torch.compile(backend=npu_backend, dynamic=False, fullgraph=True)
def opp(x, w_buffer, w_partial):#, w1, w2):
    w1 = w_buffer[:8*16].view(8, 16)
    w2 = w_buffer[8*16:].view(16, 8)
    dist.all_gather_into_tensor(w_buffer, w_partial)
    tng.scope.npu_wait_tensor(w1, w_buffer)
    h1 = torch.matmul(x, w1)
    o = torch.matmul(h1, w2)
    return o


def main():
    x = torch.randn(2, 8, device="npu", dtype=torch.bfloat16)
    wqk_buffer = torch.arange(8*16*2, device="npu", dtype=torch.bfloat16)
    wqk_partial = torch.arange(8*16*2//world_size, device="npu", dtype=torch.bfloat16)

    # wq = wqk_buffer[:8*16].view(8, 16)
    # wk = wqk_buffer[8*16:].view(16, 8)

    if local_rank == 0:
        PROFILING_PATH = "op_test"

        EXPERIMENTAL_CONFIG = torch_npu.profiler._ExperimentalConfig(
            aic_metrics=torch_npu.profiler.AiCMetrics.PipeUtilization,
            profiler_level=torch_npu.profiler.ProfilerLevel.Level1,
            l2_cache=False,
        )

        with torch_npu.profiler.profile(
            activities=[
                torch_npu.profiler.ProfilerActivity.CPU,
                torch_npu.profiler.ProfilerActivity.NPU,
            ],
            on_trace_ready=torch_npu.profiler.tensorboard_trace_handler(PROFILING_PATH),
            record_shapes=True,
            with_flops=False,
            with_modules=False,
            experimental_config=EXPERIMENTAL_CONFIG,
        ) as prof:
            o_bar = opp(x, wqk_buffer, wqk_partial)#, wq, wk)

    else:
        o_bar = opp(x, wqk_buffer, wqk_partial)#, wq, wk)

    torch.npu.synchronize()
    print_rank0(o_bar)

main()

likedislike
tangjie66
tangjie66成员
1月12日 评论:

总结:
用户场景不想将w1 w2声明放到图里面,因此这种情况下正常方式没办法避免copy的产生。
最终由用户通过from_blob的方式自定义创建tensor,使pytorch追踪不到tensor间的关系。该方式需要用户自己保证风险。

likedislike
tangjie66tangjie66成员
1月12日 issue状态由 TODO 改变为 DONE
tangjie66tangjie66成员
1月12日 关闭了 issue