已合并
fix(inductor): support Flex Attention mask_out cppwrapper #45507
fix(inductor): support Flex Attention mask_out cppwrapper #45507
已合并
Xuan Peng创建于 8月29日
共 3 个文件变更+101-4
@@ -190,10 +190,21 @@ class DeferredNpuTritonCallWrapper(DeferredTritonCallWrapper):
190 enable_simt = npu_config.is_ascend950 and (190 enable_simt = npu_config.is_ascend950 and (
191 "simt" in params["parallel_mode"] or params["force_simt_only"]191 "simt" in params["parallel_mode"] or params["force_simt_only"]
192 )192 )
193- enable_auto_blockify = not params.get("has_auto_blockify_blacklist_op", False) and triton_support_auto_blockify()193+ enable_auto_blockify = not params.get(
194+ "has_auto_blockify_blacklist_op", False
195+ ) and triton_support_auto_blockify()
196+ args_decl = wrapper.generate_args_decl(
197+ prefix,
198+ call_args,
199+ arg_types,
200+ arg_signatures,
201+ is_triton_kernel=True,
202+ force_simt_only=force_simt_only,
203+ kernel_params=params,
204+ )
194 prefix.splice(f"""205 prefix.splice(f"""
195 auto launch_call = [=]() {{206 auto launch_call = [=]() {{
196- {wrapper.generate_args_decl(prefix, call_args, arg_types, arg_signatures, True, force_simt_only)}207+ {args_decl}
197 {wrapper.generate_launch_preparation(kernel_var_name, params, enable_simt, enable_auto_blockify)}208 {wrapper.generate_launch_preparation(kernel_var_name, params, enable_simt, enable_auto_blockify)}
198 }};209 }};
199 """)210 """)
@@ -336,6 +347,7 @@ class CppWrapperNpu(CppWrapperGpu):
336 #include <acl/acl_rt.h>347 #include <acl/acl_rt.h>
337 #include <runtime/runtime/rt.h>348 #include <runtime/runtime/rt.h>
338 #include <torch_npu/csrc/core/npu/NPUStream.h>349 #include <torch_npu/csrc/core/npu/NPUStream.h>
350+ #include <torch_npu/csrc/core/npu/NPUWorkspaceAllocator.h>
339 #include <torch_npu/csrc/framework/OpCommand.h>351 #include <torch_npu/csrc/framework/OpCommand.h>
340 """352 """
341 if V.graph.aot_mode:353 if V.graph.aot_mode:
@@ -415,6 +427,7 @@ class CppWrapperNpu(CppWrapperGpu):
415 arg_signatures,427 arg_signatures,
416 is_triton_kernel=True,428 is_triton_kernel=True,
417 force_simt_only=False,429 force_simt_only=False,
430+ kernel_params: Optional[dict[str, Any]] = None,
418 ):431 ):
419 """432 """
420 Generates any declarations of args to pass into a kernel call, and then returns the arg names.433 Generates any declarations of args to pass into a kernel call, and then returns the arg names.
@@ -454,6 +467,10 @@ class CppWrapperNpu(CppWrapperGpu):
454 struct_arg_body = ""467 struct_arg_body = ""
455 468 
456 target_support_ffts = triton_support_ffts()469 target_support_ffts = triton_support_ffts()
470+ kernel_params = kernel_params or {}
471+ lock_num = int(kernel_params.get("lock_num", 0) or 0)
472+ lock_init_val = int(kernel_params.get("lock_init_val", 0) or 0)
473+ workspace_size = int(kernel_params.get("workspace_size", 0) or 0)
457 474 
458 def process_args(arg, arg_type, arg_signature=None):475 def process_args(arg, arg_type, arg_signature=None):
459 var_name = f"var_{next(self.arg_var_id)}"476 var_name = f"var_{next(self.arg_var_id)}"
@@ -515,11 +532,47 @@ class CppWrapperNpu(CppWrapperGpu):
515 }532 }
516 """533 """
517 534 
535+ workspace_str = ""
536+ if not force_simt_only and workspace_size > 0:
537+ workspace_str = f"""
538+ uint64_t workspace_size = static_cast<uint64_t>({workspace_size})
539+ * grid_0 * grid_1 * grid_2;
540+ auto workspace_tensor = at_npu::native::allocate_workspace(
541+ workspace_size, stream_);
542+ workspace_addr = const_cast<void *>(workspace_tensor.storage().data());
543+ """
544+ 
545+ sync_block_lock_str = ""
546+ if not force_simt_only and lock_num > 0:
547+ sync_block_lock_str = f"""
548+ uint64_t sync_block_lock_size = static_cast<uint64_t>({lock_num})
549+ * sizeof(int64_t);
550+ auto sync_block_lock_tensor = at_npu::native::allocate_workspace(
551+ sync_block_lock_size, stream_);
552+ sync_block_lock = const_cast<void *>(
553+ sync_block_lock_tensor.storage().data());
554+ std::vector<int64_t> sync_block_lock_init(
555+ {lock_num}, static_cast<int64_t>({lock_init_val}));
556+ ret = aclrtMemcpy(
557+ sync_block_lock,
558+ sync_block_lock_size,
559+ sync_block_lock_init.data(),
560+ sync_block_lock_size,
561+ ACL_MEMCPY_HOST_TO_DEVICE);
562+ if (ret != ACL_SUCCESS) {{
563+ throw std::runtime_error(
564+ std::string("initialize Triton sync block lock failed, 0x")
565+ + std::to_string(ret));
566+ }}
567+ """
568+ 
518 args_str = f"""569 args_str = f"""
519 aclError ret;570 aclError ret;
520 {ffts_str if target_support_ffts else ""}571 {ffts_str if target_support_ffts else ""}
521 {"void* workspace_addr = NULL;" if not force_simt_only else ""}572 {"void* workspace_addr = NULL;" if not force_simt_only else ""}
522 {"void* sync_block_lock = NULL;" if not force_simt_only else ""}573 {"void* sync_block_lock = NULL;" if not force_simt_only else ""}
574+ {workspace_str}
575+ {sync_block_lock_str}
523 struct __attribute__((packed)) {{576 struct __attribute__((packed)) {{
524 {"void* ffts_addr __attribute__((aligned(8)));" if target_support_ffts else ""}577 {"void* ffts_addr __attribute__((aligned(8)));" if target_support_ffts else ""}
525 {"void* sync_block_lock __attribute__((aligned(8)));" if not force_simt_only else ""}578 {"void* sync_block_lock __attribute__((aligned(8)));" if not force_simt_only else ""}
@@ -65,7 +65,10 @@ from torch._inductor.lowering import (
65 to_dtype,65 to_dtype,
66)66)
67from torch._inductor.runtime.runtime_utils import is_power_of_2, next_power_of_267from torch._inductor.runtime.runtime_utils import is_power_of_2, next_power_of_2
68-from torch._inductor.select_algorithm import autotune_select_algorithm68+from torch._inductor.select_algorithm import (
69+ autotune_select_algorithm,
70+ SymbolicGridFn,
71+)
69from torch.nn.attention import flex_attention as flex_attention_module72from torch.nn.attention import flex_attention as flex_attention_module
70from torch._inductor.kernel.flex_attention import (73from torch._inductor.kernel.flex_attention import (
71 flex_attention_backward_template as upstream_flex_attention_backward_template,74 flex_attention_backward_template as upstream_flex_attention_backward_template,
@@ -91,6 +94,34 @@ _LN2 = 0.6931471805599453
91_LOG2E = 1.442695040888963494_LOG2E = 1.4426950408889634
92 95 
93 96 
97+@SymbolicGridFn
98+def _symbolic_flex_attention_backward_grid(
99+ batch_size,
100+ q_heads,
101+ num_queries,
102+ d_model,
103+ kv_heads,
104+ num_key_value,
105+ meta,
106+ *,
107+ cdiv,
108+):
109+ """Backport symbolic grid support required by the C++ wrapper."""
110+ return (
111+ cdiv(num_queries, meta["BLOCK_M2"]) * (q_heads // kv_heads)
112+ + cdiv(num_key_value, meta["BLOCK_N1"]),
113+ 1,
114+ batch_size * kv_heads,
115+ )
116+ 
117+ 
118+# PyTorch 2.7.1 defines this template with a plain grid function, which cannot
119+# emit C++ expressions when its call sizes are symbolic.
120+upstream_flex_attention_backward_template.grid = (
121+ _symbolic_flex_attention_backward_grid
122+)
123+ 
124+ 
94def _tag_flex_attention_report_choices(new_choices, cfg):125def _tag_flex_attention_report_choices(new_choices, cfg):
95 """Attach tiling metadata used by NPU choice diagnostics."""126 """Attach tiling metadata used by NPU choice diagnostics."""
96 report_config = {127 report_config = {
@@ -1316,7 +1316,20 @@ class NPUCachingAutotuner(CachingAutotuner):
1316 "mix_mode": input_launcher.bin.metadata.mix_mode,1316 "mix_mode": input_launcher.bin.metadata.mix_mode,
1317 "parallel_mode": input_launcher.bin.metadata.parallel_mode,1317 "parallel_mode": input_launcher.bin.metadata.parallel_mode,
1318 "force_simt_only": input_launcher.bin.metadata.force_simt_only,1318 "force_simt_only": input_launcher.bin.metadata.force_simt_only,
1319- "has_auto_blockify_blacklist_op": getattr(input_launcher.bin.metadata, "has_auto_blockify_blacklist_op", False)1319+ "has_auto_blockify_blacklist_op": getattr(
1320+ input_launcher.bin.metadata,
1321+ "has_auto_blockify_blacklist_op",
1322+ False,
1323+ ),
1324+ "lock_num": int(
1325+ getattr(input_launcher.bin.metadata, "lock_num", 0) or 0
1326+ ),
1327+ "lock_init_val": int(
1328+ getattr(input_launcher.bin.metadata, "lock_init_val", 0) or 0
1329+ ),
1330+ "workspace_size": int(
1331+ getattr(input_launcher.bin.metadata, "workspace_size", 0) or 0
1332+ ),
1320 }1333 }
1321 enable_simt = ("simt" in params["parallel_mode"]) or params["force_simt_only"]1334 enable_simt = ("simt" in params["parallel_mode"]) or params["force_simt_only"]
1322 if npu_config.is_ascend950 and enable_simt:1335 if npu_config.is_ascend950 and enable_simt: