已合并
fix(inductor): support Flex Attention mask_out cppwrapper #45507
Xuan Peng创建于 8月29日
fix(inductor): support Flex Attention mask_out cppwrapper #45507
已合并
共 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 | ) |
| 67 | from torch._inductor.runtime.runtime_utils import is_power_of_2, next_power_of_2 | 67 | from torch._inductor.runtime.runtime_utils import is_power_of_2, next_power_of_2 |
| 68 | -from torch._inductor.select_algorithm import autotune_select_algorithm | 68 | +from torch._inductor.select_algorithm import ( |
| 69 | + autotune_select_algorithm, | ||
| 70 | + SymbolicGridFn, | ||
| 71 | +) | ||
| 69 | from torch.nn.attention import flex_attention as flex_attention_module | 72 | from torch.nn.attention import flex_attention as flex_attention_module |
| 70 | from torch._inductor.kernel.flex_attention import ( | 73 | from 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.4426950408889634 | 94 | _LOG2E = 1.4426950408889634 |
| 92 | 95 | ||
| 93 | 96 | ||
| 97 | + | ||
| 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 | + | ||
| 94 | def _tag_flex_attention_report_choices(new_choices, cfg): | 125 | def _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: |