已合并
fix_rng_state #32366
Ambi创建于 3月25日
fix_rng_state #32366
已合并
Ambi创建于 3月25日
3 个文件变更+90-0
@@ -0,0 +1,46 @@
1+import torch
2+from torch.utils.checkpoint import checkpoint
3+from torch.testing._internal.common_utils import (
4+ run_tests,
5+ instantiate_parametrized_tests,
6+)
7+from testutils import TestUtils
8+import torch_npu
9+
10+class TestDropoutWithCheckpointRecompute(TestUtils):
11+ def test_dropout_with_checkpoint_recompute(self):
12+ device = "npu"
13+
14+ def gn(x):
15+ return torch.sigmoid(torch.dropout(torch.sigmoid(x), p=0.5, train=True))
16+
17+ def fn(x):
18+ return checkpoint(
19+ gn,
20+ x,
21+ use_reentrant=False,
22+ preserve_rng_state=True,
23+ )
24+
25+ x = torch.randn(4, 4, requires_grad=True, device=device)
26+
27+ torch.manual_seed(42)
28+ eager_out = fn(x)
29+ eager_out.sum().backward()
30+ eager_grad = x.grad.clone()
31+
32+ x.grad = None
33+
34+ torch.manual_seed(42)
35+ compiled_fn = torch.compile(fn, backend="inductor")
36+ compiled_out = compiled_fn(x)
37+ compiled_out.sum().backward()
38+ compiled_grad = x.grad.clone()
39+
40+ self.assertEqual(eager_out, compiled_out)
41+ self.assertEqual(eager_grad, compiled_grad)
42+ 
43+instantiate_parametrized_tests(TestDropoutWithCheckpointRecompute)
44+ 
45+if __name__ == "__main__":
46+ run_tests()
@@ -6,6 +6,7 @@ from torch.testing._internal.common_utils import (
6)6)
7from testutils import TestUtils7from testutils import TestUtils
8import torch_npu8import torch_npu
9+import torch_npu._inductor
9 10 
10 11 
11class TestRunWithRngState(TestUtils):12class TestRunWithRngState(TestUtils):
@@ -127,11 +127,15 @@ def patch_register_philox_rand():
127def patch_register_run_and_save_rng_state_op():127def patch_register_run_and_save_rng_state_op():
128 from torch._prims import rng_prims128 from torch._prims import rng_prims
129 from torch._C import DispatchKey129 from torch._C import DispatchKey
130+ from torch._subclasses.fake_tensor import FakeTensorMode
130 131 
131 run_and_save_rng_state = getattr(132 run_and_save_rng_state = getattr(
132 rng_prims, "run_and_save_rng_state", None133 rng_prims, "run_and_save_rng_state", None
133 )134 )
134 135 
136+ if getattr(run_and_save_rng_state, "_npu_patched", False):
137+ return
138+ 
135 139 
136 @run_and_save_rng_state.py_impl(DispatchKey.PrivateUse1)140 @run_and_save_rng_state.py_impl(DispatchKey.PrivateUse1)
137 def impl_npu(op, *args, **kwargs):141 def impl_npu(op, *args, **kwargs):
@@ -143,6 +147,10 @@ def patch_register_run_and_save_rng_state_op():
143 DispatchKey.BackendSelect, None147 DispatchKey.BackendSelect, None
144 )148 )
145 149 
150+ fake_tensor_mode_impl = run_and_save_rng_state.python_key_table.get(
151+ FakeTensorMode, None
152+ )
153+ 
146 154 
147 def backend_select_with_npu(op, *args, **kwargs):155 def backend_select_with_npu(op, *args, **kwargs):
148 from torch._prims.rng_prims import get_device156 from torch._prims.rng_prims import get_device
@@ -154,14 +162,30 @@ def patch_register_run_and_save_rng_state_op():
154 162 
155 return backend_select_impl(op, *args, **kwargs)163 return backend_select_impl(op, *args, **kwargs)
156 164 
165+ 
166+ def fake_tensor_mode_with_npu(mode, op, *args, **kwargs):
167+ from torch._prims.rng_prims import get_device
168+ 
169+ device = get_device(args, kwargs)
170+ 
171+ if device == "npu":
172+ with mode:
173+ return impl_npu(op, *args, **kwargs)
174+ 
175+ return fake_tensor_mode_impl(mode, op, *args, **kwargs)
176+ 
157 run_and_save_rng_state.py_kernels[177 run_and_save_rng_state.py_kernels[
158 DispatchKey.BackendSelect178 DispatchKey.BackendSelect
159 ] = backend_select_with_npu179 ] = backend_select_with_npu
180+ run_and_save_rng_state.python_key_table[
181+ FakeTensorMode
182+ ] = fake_tensor_mode_with_npu
160 183 
161 184 
162def patch_register_run_with_rng_state_op():185def patch_register_run_with_rng_state_op():
163 from torch._prims import rng_prims186 from torch._prims import rng_prims
164 from torch._C import DispatchKey187 from torch._C import DispatchKey
188+ from torch._subclasses.fake_tensor import FakeTensorMode
165 189 
166 run_with_rng_state = getattr(190 run_with_rng_state = getattr(
167 rng_prims, "run_with_rng_state", None191 rng_prims, "run_with_rng_state", None
@@ -187,6 +211,10 @@ def patch_register_run_with_rng_state_op():
187 DispatchKey.BackendSelect, None211 DispatchKey.BackendSelect, None
188 )212 )
189 213 
214+ fake_tensor_mode_impl = run_with_rng_state.python_key_table.get(
215+ FakeTensorMode, None
216+ )
217+ 
190 218 
191 def backend_select_with_npu(rng_state, op, *args, **kwargs):219 def backend_select_with_npu(rng_state, op, *args, **kwargs):
192 from torch._prims.rng_prims import get_device220 from torch._prims.rng_prims import get_device
@@ -198,9 +226,24 @@ def patch_register_run_with_rng_state_op():
198 226 
199 return backend_select_impl(rng_state, op, *args, **kwargs)227 return backend_select_impl(rng_state, op, *args, **kwargs)
200 228 
229+ 
230+ def fake_tensor_mode_with_npu(mode, rng_state, op, *args, **kwargs):
231+ from torch._prims.rng_prims import get_device
232+ 
233+ device = get_device(args, kwargs)
234+ 
235+ if device == "npu":
236+ with mode:
237+ return op(*args, **kwargs)
238+ 
239+ return fake_tensor_mode_impl(mode, rng_state, op, *args, **kwargs)
240+ 
201 run_with_rng_state.py_kernels[241 run_with_rng_state.py_kernels[
202 DispatchKey.BackendSelect242 DispatchKey.BackendSelect
203 ] = backend_select_with_npu243 ] = backend_select_with_npu
244+ run_with_rng_state.python_key_table[
245+ FakeTensorMode
246+ ] = fake_tensor_mode_with_npu
204 247 
205 248 
206def patch_rng_prims_device():249def patch_rng_prims_device():