已合并
fix_rng_state #32366
Ambi创建于 3月25日
fix_rng_state #32366
已合并
共 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 | ) |
| 7 | from testutils import TestUtils | 7 | from testutils import TestUtils |
| 8 | import torch_npu | 8 | import torch_npu |
| 9 | +import torch_npu._inductor | ||
| 9 | 10 | ||
| 10 | 11 | ||
| 11 | class TestRunWithRngState(TestUtils): | 12 | class TestRunWithRngState(TestUtils): |
| @@ -127,11 +127,15 @@ def patch_register_philox_rand(): | |||
| 127 | def patch_register_run_and_save_rng_state_op(): | 127 | def patch_register_run_and_save_rng_state_op(): |
| 128 | from torch._prims import rng_prims | 128 | from torch._prims import rng_prims |
| 129 | from torch._C import DispatchKey | 129 | 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", None | 133 | 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 | 140 | ||
| 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, None | 147 | 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_device | 156 | 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.BackendSelect | 178 | DispatchKey.BackendSelect |
| 159 | ] = backend_select_with_npu | 179 | ] = 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 | ||
| 162 | def patch_register_run_with_rng_state_op(): | 185 | def patch_register_run_with_rng_state_op(): |
| 163 | from torch._prims import rng_prims | 186 | from torch._prims import rng_prims |
| 164 | from torch._C import DispatchKey | 187 | 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", None | 191 | rng_prims, "run_with_rng_state", None |
| @@ -187,6 +211,10 @@ def patch_register_run_with_rng_state_op(): | |||
| 187 | DispatchKey.BackendSelect, None | 211 | 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_device | 220 | 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.BackendSelect | 242 | DispatchKey.BackendSelect |
| 203 | ] = backend_select_with_npu | 243 | ] = backend_select_with_npu |
| 244 | + run_with_rng_state.python_key_table[ | ||
| 245 | + FakeTensorMode | ||
| 246 | + ] = fake_tensor_mode_with_npu | ||
| 204 | 247 | ||
| 205 | 248 | ||
| 206 | def patch_rng_prims_device(): | 249 | def patch_rng_prims_device(): |