已合并
[feat] Add test cases for verifying torch.jit.ScriptModule hook API on NPU #34411
考试请神请到孔子发现考的是英语创建于 4月25日
[feat] Add test cases for verifying torch.jit.ScriptModule hook API on NPU #34411
已合并
考试请神请到孔子发现考的是英语创建于 4月25日
1 个文件变更+489-0
@@ -5,13 +5,23 @@ Add validation cases for torch.jit APIs:
52. This file validates:52. This file validates:
6torch.jit.onednn_fusion_enabled6torch.jit.onednn_fusion_enabled
7torch.jit.enable_onednn_fusion7torch.jit.enable_onednn_fusion
8+torch.jit.ScriptModule.register_full_backward_hook,
9+torch.jit.ScriptModule.register_full_backward_pre_hook,
10+torch.jit.ScriptModule.register_load_state_dict_pre_hook,
11+torch.jit.ScriptModule.register_load_state_dict_post_hook,
12+torch.jit.ScriptModule.register_state_dict_pre_hook,
13+torch.jit.ScriptModule.register_state_dict_post_hook
8(extendable)14(extendable)
9"""15"""
10 16 
11import torch17import torch
18+import torch.nn as nn
12from torch.testing._internal.common_utils import run_tests, TestCase19from torch.testing._internal.common_utils import run_tests, TestCase
13 20 
14 21 
22+device_type = acc.type if (acc := torch.accelerator.current_accelerator()) else "cpu"
23+ 
24+ 
15class TestOneDNNJitAPI(TestCase):25class TestOneDNNJitAPI(TestCase):
16 def setUp(self):26 def setUp(self):
17 self.original_state = torch.jit.onednn_fusion_enabled()27 self.original_state = torch.jit.onednn_fusion_enabled()
@@ -33,5 +43,484 @@ class TestOneDNNJitAPI(TestCase):
33 self.assertEqual(torch.jit.onednn_fusion_enabled(), False)43 self.assertEqual(torch.jit.onednn_fusion_enabled(), False)
34 44 
35 45 
46+class TestScriptModuleHooks(TestCase):
47+ # register_full_backward_hook: not working for ScriptModule, should raise error
48+ 
49+ def test_register_full_backward_hook_raises(self):
50+ class M(torch.jit.ScriptModule):
51+ def __init__(self):
52+ super().__init__()
53+ self.linear = nn.Linear(2, 2)
54+ 
55+ def forward(self, x):
56+ return self.linear(x)
57+ 
58+ model = M().to(device_type)
59+ with self.assertRaises(RuntimeError):
60+ model.register_full_backward_hook(
61+ lambda module, grad_input, grad_output: grad_input
62+ )
63+ 
64+ # register_full_backward_pre_hook
65+ 
66+ def test_register_full_backward_pre_hook_called(self):
67+ class M(torch.jit.ScriptModule):
68+ def __init__(self):
69+ super().__init__()
70+ self.linear = nn.Linear(2, 2)
71+ 
72+ def forward(self, x):
73+ return self.linear(x)
74+ 
75+ model = M().to(device_type)
76+ called = []
77+ handle = model.register_full_backward_pre_hook(
78+ lambda module, grad_output: (called.append(True), grad_output)[1]
79+ )
80+ x = torch.randn(2, 2, requires_grad=True).to(device_type)
81+ model(x).sum().backward()
82+ self.assertTrue(called)
83+ handle.remove()
84+ 
85+ def test_register_full_backward_pre_hook_modify_grad(self):
86+ class M(torch.jit.ScriptModule):
87+ def __init__(self):
88+ super().__init__()
89+ self.linear = nn.Linear(2, 2)
90+ 
91+ def forward(self, x):
92+ return self.linear(x)
93+ 
94+ model = M().to(device_type)
95+ model.register_full_backward_pre_hook(
96+ lambda module, grad_output: (grad_output[0] * 0,)
97+ )
98+ x = torch.ones(2, device=device_type, requires_grad=True)
99+ model(x).sum().backward()
100+ self.assertEqual(x.grad, torch.zeros(2, device=device_type))
101+ 
102+ def test_register_full_backward_pre_hook_prepend(self):
103+ class M(torch.jit.ScriptModule):
104+ def __init__(self):
105+ super().__init__()
106+ self.linear = nn.Linear(2, 2)
107+ 
108+ def forward(self, x):
109+ return self.linear(x)
110+ 
111+ model = M().to(device_type)
112+ order = []
113+ model.register_full_backward_pre_hook(
114+ lambda module, grad_output: (order.append(1), grad_output)[1]
115+ )
116+ model.register_full_backward_pre_hook(
117+ lambda module, grad_output: (order.append(2), grad_output)[1],
118+ prepend=True,
119+ )
120+ x = torch.randn(2, 2, requires_grad=True).to(device_type)
121+ model(x).sum().backward()
122+ self.assertEqual(order, [2, 1])
123+ 
124+ def test_register_full_backward_pre_hook_remove(self):
125+ class M(torch.jit.ScriptModule):
126+ def __init__(self):
127+ super().__init__()
128+ self.linear = nn.Linear(2, 2)
129+ 
130+ def forward(self, x):
131+ return self.linear(x)
132+ 
133+ model = M().to(device_type)
134+ called = []
135+ handle = model.register_full_backward_pre_hook(
136+ lambda module, grad_output: (called.append(True), grad_output)[1]
137+ )
138+ x = torch.randn(2, 2, requires_grad=True).to(device_type)
139+ model(x).sum().backward()
140+ self.assertEqual(len(called), 1)
141+ handle.remove()
142+ model(x).sum().backward()
143+ self.assertEqual(len(called), 1)
144+ 
145+ # register_load_state_dict_pre_hook & register_load_state_dict_post_hook
146+ 
147+ def test_load_state_dict_pre_hook_fires_before_module_and_post_hook(self):
148+ class M(torch.jit.ScriptModule):
149+ def __init__(self):
150+ super().__init__()
151+ self.linear = nn.Linear(2, 2)
152+ 
153+ def forward(self, x):
154+ return self.linear(x)
155+ 
156+ order = []
157+ 
158+ def post_hook(module, incompatible_keys):
159+ order.append("post")
160+ 
161+ def pre_hook(
162+ module,
163+ state_dict,
164+ prefix,
165+ local_metadata,
166+ strict,
167+ missing_keys,
168+ unexpected_keys,
169+ error_msgs,
170+ ):
171+ order.append("pre")
172+ 
173+ model = M().to(device_type)
174+ model.register_load_state_dict_post_hook(post_hook)
175+ model.register_load_state_dict_pre_hook(pre_hook)
176+ model.load_state_dict(model.state_dict())
177+ self.assertEqual(order, ["pre", "post"])
178+ 
179+ order2 = []
180+ 
181+ def pre_hook_modify(
182+ module,
183+ state_dict,
184+ prefix,
185+ local_metadata,
186+ strict,
187+ missing_keys,
188+ unexpected_keys,
189+ error_msgs,
190+ ):
191+ order2.append("pre")
192+ for key in list(state_dict.keys()):
193+ state_dict[key] = torch.zeros_like(state_dict[key]).to(device_type)
194+ 
195+ def post_hook_check(module, incompatible_keys):
196+ order2.append("post")
197+ self.assertEqual(
198+ module.linear.weight,
199+ torch.zeros_like(module.linear.weight).to(device_type),
200+ )
201+ 
202+ model2 = M().to(device_type)
203+ model2.register_load_state_dict_post_hook(post_hook_check)
204+ model2.register_load_state_dict_pre_hook(pre_hook_modify)
205+ model2.load_state_dict(model.state_dict())
206+ self.assertEqual(order2, ["pre", "post"])
207+ 
208+ def test_register_load_state_dict_pre_hook_called(self):
209+ class M(torch.jit.ScriptModule):
210+ def __init__(self):
211+ super().__init__()
212+ self.linear = nn.Linear(2, 2)
213+ 
214+ def forward(self, x):
215+ return self.linear(x)
216+ 
217+ model = M().to(device_type)
218+ called = []
219+ 
220+ def hook(
221+ module,
222+ state_dict,
223+ prefix,
224+ local_metadata,
225+ strict,
226+ missing_keys,
227+ unexpected_keys,
228+ error_msgs,
229+ ):
230+ called.append(prefix)
231+ 
232+ model.register_load_state_dict_pre_hook(hook)
233+ model.load_state_dict(model.state_dict())
234+ self.assertEqual(called, [""])
235+ 
236+ def test_register_load_state_dict_pre_hook_with_module(self):
237+ class M(torch.jit.ScriptModule):
238+ def __init__(self):
239+ super().__init__()
240+ self.linear = nn.Linear(2, 2)
241+ 
242+ def forward(self, x):
243+ return self.linear(x)
244+ 
245+ model = M().to(device_type)
246+ received_module = []
247+ 
248+ def hook(
249+ module,
250+ state_dict,
251+ prefix,
252+ local_metadata,
253+ strict,
254+ missing_keys,
255+ unexpected_keys,
256+ error_msgs,
257+ ):
258+ received_module.append(module)
259+ 
260+ model.register_load_state_dict_pre_hook(hook)
261+ model.load_state_dict(model.state_dict())
262+ self.assertIs(received_module[0], model)
263+ 
264+ def test_register_load_state_dict_pre_hook_remove(self):
265+ class M(torch.jit.ScriptModule):
266+ def __init__(self):
267+ super().__init__()
268+ self.linear = nn.Linear(2, 2)
269+ 
270+ def forward(self, x):
271+ return self.linear(x)
272+ 
273+ model = M().to(device_type)
274+ called = []
275+ 
276+ def hook(
277+ module,
278+ state_dict,
279+ prefix,
280+ local_metadata,
281+ strict,
282+ missing_keys,
283+ unexpected_keys,
284+ error_msgs,
285+ ):
286+ called.append(True)
287+ 
288+ handle = model.register_load_state_dict_pre_hook(hook)
289+ model.load_state_dict(model.state_dict())
290+ self.assertEqual(len(called), 1)
291+ handle.remove()
292+ model.load_state_dict(model.state_dict())
293+ self.assertEqual(len(called), 1)
294+ 
295+ # register_load_state_dict_post_hook
296+ 
297+ def test_register_load_state_dict_post_hook_called(self):
298+ class M(torch.jit.ScriptModule):
299+ def __init__(self):
300+ super().__init__()
301+ self.linear = nn.Linear(2, 2)
302+ 
303+ def forward(self, x):
304+ return self.linear(x)
305+ 
306+ model = M().to(device_type)
307+ called = []
308+ 
309+ def hook(module, incompatible_keys):
310+ called.append(True)
311+ 
312+ handle = model.register_load_state_dict_post_hook(hook)
313+ model.load_state_dict(model.state_dict())
314+ self.assertTrue(called)
315+ handle.remove()
316+ 
317+ def test_register_load_state_dict_post_hook_with_module(self):
318+ class M(torch.jit.ScriptModule):
319+ def __init__(self):
320+ super().__init__()
321+ self.linear = nn.Linear(2, 2)
322+ 
323+ def forward(self, x):
324+ return self.linear(x)
325+ 
326+ model = M().to(device_type)
327+ received_module = []
328+ 
329+ def hook(module, incompatible_keys):
330+ received_module.append(module)
331+ 
332+ handle = model.register_load_state_dict_post_hook(hook)
333+ model.load_state_dict(model.state_dict())
334+ self.assertIs(received_module[0], model)
335+ handle.remove()
336+ 
337+ def test_register_load_state_dict_post_hook_remove(self):
338+ class M(torch.jit.ScriptModule):
339+ def __init__(self):
340+ super().__init__()
341+ self.linear = nn.Linear(2, 2)
342+ 
343+ def forward(self, x):
344+ return self.linear(x)
345+ 
346+ model = M().to(device_type)
347+ called = []
348+ 
349+ def hook(module, incompatible_keys):
350+ called.append(True)
351+ 
352+ handle = model.register_load_state_dict_post_hook(hook)
353+ model.load_state_dict(model.state_dict())
354+ self.assertEqual(len(called), 1)
355+ handle.remove()
356+ model.load_state_dict(model.state_dict())
357+ self.assertEqual(len(called), 1)
358+ 
359+ # register_state_dict_pre_hook & register_state_dict_post_hook
360+ 
361+ def test_state_dict_pre_hook_fires_before_module_and_post_hook(self):
362+ class M(torch.jit.ScriptModule):
363+ def __init__(self):
364+ super().__init__()
365+ self.linear = nn.Linear(2, 2)
366+ 
367+ def forward(self, x):
368+ return self.linear(x)
369+ 
370+ order = []
371+ 
372+ def post_hook(module, state_dict, prefix, local_metadata):
373+ order.append("post")
374+ 
375+ def pre_hook(module, prefix, keep_vars):
376+ order.append("pre")
377+ 
378+ model = M().to(device_type)
379+ model.register_state_dict_post_hook(post_hook)
380+ model.register_state_dict_pre_hook(pre_hook)
381+ model.state_dict()
382+ self.assertEqual(order, ["pre", "post"])
383+ 
384+ order2 = []
385+ pre_hook_called = [False]
386+ 
387+ def pre_hook_flag(module, prefix, keep_vars):
388+ order2.append("pre")
389+ pre_hook_called[0] = True
390+ 
391+ def post_hook_check(module, state_dict, prefix, local_metadata):
392+ order2.append("post")
393+ self.assertTrue(pre_hook_called[0])
394+ self.assertTrue(len(state_dict) > 0)
395+ 
396+ model2 = M().to(device_type)
397+ model2.register_state_dict_post_hook(post_hook_check)
398+ model2.register_state_dict_pre_hook(pre_hook_flag)
399+ model2.state_dict()
400+ self.assertEqual(order2, ["pre", "post"])
401+ 
402+ def test_register_state_dict_pre_hook_called(self):
403+ class M(torch.jit.ScriptModule):
404+ def __init__(self):
405+ super().__init__()
406+ self.linear = nn.Linear(2, 2)
407+ 
408+ def forward(self, x):
409+ return self.linear(x)
410+ 
411+ model = M().to(device_type)
412+ called = []
413+ 
414+ def hook(module, prefix, keep_vars):
415+ called.append(prefix)
416+ 
417+ model.register_state_dict_pre_hook(hook)
418+ model.state_dict()
419+ self.assertEqual(called, [""])
420+ 
421+ def test_register_state_dict_pre_hook_with_module(self):
422+ class M(torch.jit.ScriptModule):
423+ def __init__(self):
424+ super().__init__()
425+ self.linear = nn.Linear(2, 2)
426+ 
427+ def forward(self, x):
428+ return self.linear(x)
429+ 
430+ model = M().to(device_type)
431+ received_module = []
432+ 
433+ def hook(module, prefix, keep_vars):
434+ received_module.append(module)
435+ 
436+ model.register_state_dict_pre_hook(hook)
437+ model.state_dict()
438+ self.assertIs(received_module[0], model)
439+ 
440+ def test_register_state_dict_pre_hook_remove(self):
441+ class M(torch.jit.ScriptModule):
442+ def __init__(self):
443+ super().__init__()
444+ self.linear = nn.Linear(2, 2)
445+ 
446+ def forward(self, x):
447+ return self.linear(x)
448+ 
449+ model = M().to(device_type)
450+ called = []
451+ 
452+ def hook(module, prefix, keep_vars):
453+ called.append(True)
454+ 
455+ handle = model.register_state_dict_pre_hook(hook)
456+ model.state_dict()
457+ self.assertEqual(len(called), 1)
458+ handle.remove()
459+ model.state_dict()
460+ self.assertEqual(len(called), 1)
461+ 
462+ # register_state_dict_post_hook
463+ 
464+ def test_register_state_dict_post_hook_called(self):
465+ class M(torch.jit.ScriptModule):
466+ def __init__(self):
467+ super().__init__()
468+ self.linear = nn.Linear(2, 2)
469+ 
470+ def forward(self, x):
471+ return self.linear(x)
472+ 
473+ model = M().to(device_type)
474+ called = []
475+ 
476+ def hook(module, state_dict, prefix, local_metadata):
477+ called.append(prefix)
478+ 
479+ model.register_state_dict_post_hook(hook)
480+ model.state_dict()
481+ self.assertEqual(called, [""])
482+ 
483+ def test_register_state_dict_post_hook_with_module(self):
484+ class M(torch.jit.ScriptModule):
485+ def __init__(self):
486+ super().__init__()
487+ self.linear = nn.Linear(2, 2)
488+ 
489+ def forward(self, x):
490+ return self.linear(x)
491+ 
492+ model = M().to(device_type)
493+ received_module = []
494+ 
495+ def hook(module, state_dict, prefix, local_metadata):
496+ received_module.append(module)
497+ 
498+ model.register_state_dict_post_hook(hook)
499+ model.state_dict()
500+ self.assertIs(received_module[0], model)
501+ 
502+ def test_register_state_dict_post_hook_remove(self):
503+ class M(torch.jit.ScriptModule):
504+ def __init__(self):
505+ super().__init__()
506+ self.linear = nn.Linear(2, 2)
507+ 
508+ def forward(self, x):
509+ return self.linear(x)
510+ 
511+ model = M().to(device_type)
512+ called = []
513+ 
514+ def hook(module, state_dict, prefix, local_metadata):
515+ called.append(True)
516+ 
517+ handle = model.register_state_dict_post_hook(hook)
518+ model.state_dict()
519+ self.assertEqual(len(called), 1)
520+ handle.remove()
521+ model.state_dict()
522+ self.assertEqual(len(called), 1)
523+ 
524+ 
36if __name__ == "__main__":525if __name__ == "__main__":
37 run_tests()526 run_tests()