您好,可以尝试将
wq = wqk_buffer[:816].view(8, 16)
wk = wqk_buffer[816:].view(16, 8)
改为
wq = wqk_buffer[:816].view(8, 16).clone()
wk = wqk_buffer[816:].view(16, 8).clone()
再看下图中有没有多余的viewcopy和strided等操作
以下是对应的解释:
-
为什么原来的写法有asstrided:
以wq为例,原来的wq只是给wqk_buffer的前128个元素贴了个新标签,还是共享的一块内存,它的步长需要根据wqk_buffer的步长进行推导
而npu需要在编译阶段,通过asstrided将布局规则写死,在执行时才能知道从wqk_buffer中怎么取数据。 -
为什么原来的写法中有viewcopy:
all_gather_into_tensor执行的时候会往wqk_buffer中写新数据,wq和wqk_buffer共享一块内存
想让wq能正确读取更新的数据,必须将输出的通信结果的内存布局通过viewcopy变成wq能识别的格式 -
为什么新写法没有这些算子:
如果wq和wk不依赖其他张量,如wqk_buffer,它的内存是独立的,编译框架无需追踪两者的布局依赖


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


但是我这个场景里,需要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()


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


在提交问题之前,请通过搜索现有和历史问题确保该问题尚未被提出并解决。
您的环境信息
-- 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, 图模式会引入额外的操作
图结构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()