已合并
Support registering functions on AutogradPrivateUse1 via YAML dispatch configuration #25031
shaoyf创建于 2025年9月19日
Support registering functions on AutogradPrivateUse1 via YAML dispatch configuration #25031
已合并
共 3 个文件变更+31-13
| @@ -9,7 +9,7 @@ from torchgen.gen import (parse_tags_yaml, FileManager, cpp_string, error_check_ | |||
| 9 | from torchgen.model import (BackendIndex, DispatchKey, Variant, | 9 | from torchgen.model import (BackendIndex, DispatchKey, Variant, |
| 10 | NativeFunction, OperatorName, BackendMetadata, TensorOptionsArguments) | 10 | NativeFunction, OperatorName, BackendMetadata, TensorOptionsArguments) |
| 11 | from torchgen.utils import concatMap, mapMaybe | 11 | from torchgen.utils import concatMap, mapMaybe |
| 12 | -from torchgen.context import with_native_function, native_function_manager, method_with_native_function | 12 | +from torchgen.context import with_native_function, native_function_manager, method_with_native_function, with_native_function_and_index |
| 13 | from torchgen.api.types import DispatcherSignature | 13 | from torchgen.api.types import DispatcherSignature |
| 14 | from torchgen.api import cpp | 14 | from torchgen.api import cpp |
| 15 | from torchgen.dest.register_dispatch_key import RegisterDispatchKey | 15 | from torchgen.dest.register_dispatch_key import RegisterDispatchKey |
| @@ -63,7 +63,7 @@ def parse_custom_yaml(custom_path: str, tag_path: str) -> ParsedYaml: | |||
| 63 | error_check_native_functions(rs) | 63 | error_check_native_functions(rs) |
| 64 | # Default dict is to prevent the codegen from barfing when we have a dispatch key that has no kernels yet. | 64 | # Default dict is to prevent the codegen from barfing when we have a dispatch key that has no kernels yet. |
| 65 | indices: Dict[DispatchKey, BackendIndex] = defaultdict(lambda: BackendIndex( | 65 | indices: Dict[DispatchKey, BackendIndex] = defaultdict(lambda: BackendIndex( |
| 66 | - dispatch_key=DispatchKey.Undefined, use_out_as_primary=True, external=False, index={})) | 66 | + dispatch_key=DispatchKey.Undefined, use_out_as_primary=True, external=False, device_guard=False, index={})) |
| 67 | for k, v in bs.items(): | 67 | for k, v in bs.items(): |
| 68 | # All structured in-tree operators are implemented in terms of their out operator. | 68 | # All structured in-tree operators are implemented in terms of their out operator. |
| 69 | indices[k] = BackendIndex(dispatch_key=k, | 69 | indices[k] = BackendIndex(dispatch_key=k, |
| @@ -241,8 +241,11 @@ class RegisterCustomSchema: | |||
| 241 | return f'{maybe_tags}m.def({cpp_string(func_schema)}{tag_index});\n' | 241 | return f'{maybe_tags}m.def({cpp_string(func_schema)}{tag_index});\n' |
| 242 | 242 | ||
| 243 | 243 | ||
| 244 | -@with_native_function | 244 | +@with_native_function_and_index |
| 245 | -def compute_register_impl(f: NativeFunction): | 245 | +def compute_register_impl(f: NativeFunction, backend_index): |
| 246 | + if (backend_index is not None) and (backend_index.get_kernel(f) is None): | ||
| 247 | + return [] | ||
| 248 | + | ||
| 246 | if f.has_composite_explicit_autograd_kernel: | 249 | if f.has_composite_explicit_autograd_kernel: |
| 247 | return [] | 250 | return [] |
| 248 | else: | 251 | else: |
| @@ -250,7 +253,7 @@ def compute_register_impl(f: NativeFunction): | |||
| 250 | return [f'm.impl("{f.func.name}", TORCH_FN(at_npu::native::{name}));\n'] | 253 | return [f'm.impl("{f.func.name}", TORCH_FN(at_npu::native::{name}));\n'] |
| 251 | 254 | ||
| 252 | 255 | ||
| 253 | -def gen_custom_trace(fm: FileManager, custom_trace_functions: Sequence[NativeFunction]): | 256 | +def gen_custom_trace(fm: FileManager, custom_trace_functions: Sequence[NativeFunction], custom_backend_indices): |
| 254 | 257 | ||
| 255 | fm.write_with_template(f'CustomRegisterSchema.cpp', 'CustomRegisterSchema.cpp', lambda: { | 258 | fm.write_with_template(f'CustomRegisterSchema.cpp', 'CustomRegisterSchema.cpp', lambda: { |
| 256 | 'custom_op_definitions': list(concatMap( | 259 | 'custom_op_definitions': list(concatMap( |
| @@ -262,7 +265,11 @@ def gen_custom_trace(fm: FileManager, custom_trace_functions: Sequence[NativeFun | |||
| 262 | custom_trace_functions | 265 | custom_trace_functions |
| 263 | )), | 266 | )), |
| 264 | 'custom_impl_registrations': list(concatMap( | 267 | 'custom_impl_registrations': list(concatMap( |
| 265 | - lambda f: compute_register_impl(f), | 268 | + lambda f: compute_register_impl(f, None), |
| 269 | + custom_trace_functions | ||
| 270 | + )), | ||
| 271 | + 'custom_autograd_impl_registrations': list(concatMap( | ||
| 272 | + lambda f: compute_register_impl(f, custom_backend_indices[DispatchKey.AutogradPrivateUse1]), | ||
| 266 | custom_trace_functions | 273 | custom_trace_functions |
| 267 | )), | 274 | )), |
| 268 | }) | 275 | }) |
| @@ -23,6 +23,7 @@ from collections import namedtuple, Counter, defaultdict | |||
| 23 | from typing import List, Dict, Union, Sequence, Optional, Set, Callable | 23 | from typing import List, Dict, Union, Sequence, Optional, Set, Callable |
| 24 | import yaml | 24 | import yaml |
| 25 | 25 | ||
| 26 | +import torchgen | ||
| 26 | from torchgen.code_template import CodeTemplate | 27 | from torchgen.code_template import CodeTemplate |
| 27 | from torchgen.gen import (parse_tags_yaml, FileManager, parse_native_yaml, | 28 | from torchgen.gen import (parse_tags_yaml, FileManager, parse_native_yaml, |
| 28 | get_grouped_native_functions, error_check_native_functions) | 29 | get_grouped_native_functions, error_check_native_functions) |
| @@ -48,6 +49,8 @@ from codegen.utils import (get_torchgen_dir, rename_privateuse1_dispatch_key, ge | |||
| 48 | from codegen.custom_functions import (parse_custom_yaml, gen_custom_trace, gen_custom_ops_patch, | 49 | from codegen.custom_functions import (parse_custom_yaml, gen_custom_trace, gen_custom_ops_patch, |
| 49 | gen_custom_functions_dispatch) | 50 | gen_custom_functions_dispatch) |
| 50 | 51 | ||
| 52 | +torchgen.model.dispatch_keys.append(torchgen.model.DispatchKey.AutogradPrivateUse1) | ||
| 53 | + | ||
| 51 | 54 | ||
| 52 | # Create backend_indices map for func retrieval with the key of each func we supported. | 55 | # Create backend_indices map for func retrieval with the key of each func we supported. |
| 53 | def create_backend_index(backend_ops: List[str], | 56 | def create_backend_index(backend_ops: List[str], |
| @@ -300,8 +303,7 @@ the behavior of autograd for some operators on your backend. However "Autograd{b | |||
| 300 | autograd_key, native_functions_map, cpp_namespace) | 303 | autograd_key, native_functions_map, cpp_namespace) |
| 301 | opapi_autograd_idx = create_backend_index([op for op in supported_autograd if is_opapi(op)], | 304 | opapi_autograd_idx = create_backend_index([op for op in supported_autograd if is_opapi(op)], |
| 302 | symint_set, autograd_key, native_functions_map, cpp_namespace) | 305 | symint_set, autograd_key, native_functions_map, cpp_namespace) |
| 303 | - if autograd_key in backend_indices: | 306 | + |
| 304 | - raise KeyError("autograd_key should not be in backend_indices.") | ||
| 305 | backend_indices[autograd_key] = autograd_idx | 307 | backend_indices[autograd_key] = autograd_idx |
| 306 | backend_indices[str(autograd_key) + opapi_key] = opapi_autograd_idx | 308 | backend_indices[str(autograd_key) + opapi_key] = opapi_autograd_idx |
| 307 | 309 | ||
| @@ -732,15 +734,15 @@ def run(source_yaml: str, output_dir: str, dry_run: bool, | |||
| 732 | register_dispatch_key_func=dest.RegisterDispatchKey, | 734 | register_dispatch_key_func=dest.RegisterDispatchKey, |
| 733 | ) | 735 | ) |
| 734 | 736 | ||
| 735 | - gen_quantize_register(fm, backend_indices["NPUQuantize"]) | 737 | + gen_quantize_register(fm, backend_indices=backend_indices["NPUQuantize"]) |
| 736 | 738 | ||
| 737 | pta_template_dir = os.path.join(pathlib.Path(__file__).parent.absolute(), "templates") | 739 | pta_template_dir = os.path.join(pathlib.Path(__file__).parent.absolute(), "templates") |
| 738 | fm = FileManager(install_dir=output_dir, template_dir=pta_template_dir, dry_run=dry_run) | 740 | fm = FileManager(install_dir=output_dir, template_dir=pta_template_dir, dry_run=dry_run) |
| 739 | 741 | ||
| 740 | - custom_functions = parse_custom_yaml(source_yaml, tags_yaml_path).native_functions | 742 | + custom_functions, custom_backend_indices = parse_custom_yaml(source_yaml, tags_yaml_path) |
| 741 | grouped_custom_functions = get_grouped_native_functions_optional_out(custom_functions) | 743 | grouped_custom_functions = get_grouped_native_functions_optional_out(custom_functions) |
| 742 | gen_functionalization(fm, selector, grouped_custom_functions) | 744 | gen_functionalization(fm, selector, grouped_custom_functions) |
| 743 | - gen_custom_trace(fm, custom_functions) | 745 | + gen_custom_trace(fm, custom_functions, custom_backend_indices) |
| 744 | gen_custom_functions_dispatch(fm, custom_functions) | 746 | gen_custom_functions_dispatch(fm, custom_functions) |
| 745 | 747 | ||
| 746 | gen_foreach_register(fm, | 748 | gen_foreach_register(fm, |
| @@ -43,7 +43,7 @@ namespace { | |||
| 43 | 43 | ||
| 44 | TORCH_LIBRARY(npu, m) { | 44 | TORCH_LIBRARY(npu, m) { |
| 45 | 45 | ||
| 46 | - ${custom_schema_registrations} | 46 | + ${custom_schema_registrations} |
| 47 | } | 47 | } |
| 48 | 48 | ||
| 49 | } // anonymous namespace | 49 | } // anonymous namespace |
| @@ -52,7 +52,16 @@ namespace { | |||
| 52 | 52 | ||
| 53 | TORCH_LIBRARY_IMPL(npu, PrivateUse1, m) { | 53 | TORCH_LIBRARY_IMPL(npu, PrivateUse1, m) { |
| 54 | 54 | ||
| 55 | - ${custom_impl_registrations} | 55 | + ${custom_impl_registrations} |
| 56 | +} | ||
| 57 | + | ||
| 58 | +} // anonymous namespace | ||
| 59 | + | ||
| 60 | +namespace { | ||
| 61 | + | ||
| 62 | +TORCH_LIBRARY_IMPL(npu, AutogradPrivateUse1, m) { | ||
| 63 | + | ||
| 64 | + ${custom_autograd_impl_registrations} | ||
| 56 | } | 65 | } |
| 57 | 66 | ||
| 58 | } // anonymous namespace | 67 | } // anonymous namespace |