已合并
Support registering functions on AutogradPrivateUse1 via YAML dispatch configuration #25031
shaoyf创建于 2025年9月19日
Support registering functions on AutogradPrivateUse1 via YAML dispatch configuration #25031
已合并
shaoyf创建于 2025年9月19日
共 3 个文件变更+31-13
@@ -9,7 +9,7 @@ from torchgen.gen import (parse_tags_yaml, FileManager, cpp_string, error_check_
9from torchgen.model import (BackendIndex, DispatchKey, Variant,9from torchgen.model import (BackendIndex, DispatchKey, Variant,
10 NativeFunction, OperatorName, BackendMetadata, TensorOptionsArguments)10 NativeFunction, OperatorName, BackendMetadata, TensorOptionsArguments)
11from torchgen.utils import concatMap, mapMaybe11from torchgen.utils import concatMap, mapMaybe
12-from torchgen.context import with_native_function, native_function_manager, method_with_native_function12+from torchgen.context import with_native_function, native_function_manager, method_with_native_function, with_native_function_and_index
13from torchgen.api.types import DispatcherSignature13from torchgen.api.types import DispatcherSignature
14from torchgen.api import cpp14from torchgen.api import cpp
15from torchgen.dest.register_dispatch_key import RegisterDispatchKey15from 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_function244+@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_functions265 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_functions273 custom_trace_functions
267 )),274 )),
268 })275 })
@@ -23,6 +23,7 @@ from collections import namedtuple, Counter, defaultdict
23from typing import List, Dict, Union, Sequence, Optional, Set, Callable23from typing import List, Dict, Union, Sequence, Optional, Set, Callable
24import yaml24import yaml
25 25 
26+import torchgen
26from torchgen.code_template import CodeTemplate27from torchgen.code_template import CodeTemplate
27from torchgen.gen import (parse_tags_yaml, FileManager, parse_native_yaml,28from 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
48from codegen.custom_functions import (parse_custom_yaml, gen_custom_trace, gen_custom_ops_patch,49from 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.
53def create_backend_index(backend_ops: List[str],56def 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_idx307 backend_indices[autograd_key] = autograd_idx
306 backend_indices[str(autograd_key) + opapi_key] = opapi_autograd_idx308 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_functions742+ 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 
44TORCH_LIBRARY(npu, m) {44TORCH_LIBRARY(npu, m) {
45 45 
46- ${custom_schema_registrations}46+ ${custom_schema_registrations}
47}47}
48 48 
49} // anonymous namespace49} // anonymous namespace
@@ -52,7 +52,16 @@ namespace {
52 52 
53TORCH_LIBRARY_IMPL(npu, PrivateUse1, m) {53TORCH_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 namespace67} // anonymous namespace