已合并
Fix UT from test_register_sharding #29157
zhangguoguang创建于 1月6日
Fix UT from test_register_sharding #29157
已合并
zhangguoguang创建于 1月6日
3 个文件变更+1060-1052
@@ -223,6 +223,733 @@ class TestMathOps(NPUDTensorTestBase):
223 test_placement_comb([placement], [placement], [placement])223 test_placement_comb([placement], [placement], [placement])
224 224 
225 225 
226+class TestConv2d(NPUDTensorTestBase):
227+ @SupportedDevices(['Ascend910B'])
228+ @skipIfUnsupportMultiNPU(2)
229+ @with_comms
230+ def test_torch_npu_npu_conv2d_replicate(self):
231+ mesh = self.build_device_mesh()
232+ 
233+ input_tensor = torch.randn(3, 3, 224, 224, device="npu", requires_grad=True)
234+ weight_tensor = torch.randn(64, 3, 3, 3, device="npu", requires_grad=True)
235+ 
236+ input_dtensor = distribute_tensor(input_tensor, mesh, [Shard(0)])
237+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
238+ 
239+ bias = torch.randn(64, device="npu", requires_grad=True)
240+ d_bias = distribute_tensor(bias, mesh, [Replicate()])
241+ 
242+ 
243+ stride = (1, 1)
244+ padding = (1, 1)
245+ dilation = (1, 1)
246+ groups = 1
247+ 
248+ output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, d_bias, stride, padding, dilation, groups)
249+ output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
250+ 
251+ self.assertEqual(output_dtensor.full_tensor(), output_tensor)
252+ 
253+ @SupportedDevices(['Ascend910B'])
254+ @skipIfUnsupportMultiNPU(2)
255+ @with_comms
256+ def test_torch_npu_npu_conv2d_weight_shard0(self):
257+ mesh = self.build_device_mesh()
258+ 
259+ input_tensor = torch.randn(3, 3, 224, 224, device="npu", requires_grad=True)
260+ weight_tensor = torch.randn(64, 3, 3, 3, device="npu", requires_grad=True)
261+ 
262+ input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
263+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Shard(0)])
264+ 
265+ bias = torch.randn(64, device="npu", requires_grad=True)
266+ d_bias = distribute_tensor(bias, mesh, [Shard(0)])
267+ 
268+ 
269+ stride = (1, 1)
270+ padding = (1, 1)
271+ dilation = (1, 1)
272+ groups = 1
273+ 
274+ output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, d_bias, stride, padding, dilation, groups)
275+ output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
276+ 
277+ self.assertEqual(output_dtensor.full_tensor(), output_tensor)
278+ 
279+ @SupportedDevices(['Ascend910B'])
280+ @skipIfUnsupportMultiNPU(2)
281+ @with_comms
282+ def test_torch_npu_npu_conv2d_input_shard1(self):
283+ mesh = self.build_device_mesh()
284+ 
285+ input_tensor = torch.randn(8, 4, 224, 224, device="npu", requires_grad=True)
286+ weight_tensor = torch.randn(64, 4, 3, 3, device="npu", requires_grad=True)
287+ 
288+ input_dtensor = distribute_tensor(input_tensor, mesh, [Shard(1)])
289+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Shard(1)])
290+ 
291+ bias = torch.randn(64, device="npu", requires_grad=True)
292+ d_bias = distribute_tensor(bias, mesh, [Replicate()])
293+ 
294+ stride = (1, 1)
295+ padding = (1, 1)
296+ dilation = (1, 1)
297+ groups = 1
298+ 
299+ output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, d_bias, stride, padding, dilation, groups)
300+ output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
301+ 
302+ self.assertEqual(output_dtensor.full_tensor(), output_tensor)
303+ 
304+ @SupportedDevices(['Ascend910B'])
305+ @skipIfUnsupportMultiNPU(2)
306+ @with_comms
307+ def test_torch_npu_npu_conv2d_bias_is_None_replicate(self):
308+ mesh = self.build_device_mesh()
309+ 
310+ input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
311+ weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
312+ 
313+ input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
314+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
315+ 
316+ bias = None
317+ 
318+ stride = (1, 1)
319+ padding = (1, 1)
320+ dilation = (1, 1)
321+ groups = 1
322+ 
323+ output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
324+ output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, bias, stride, padding, dilation, groups)
325+ self.assertEqual(output_dtensor.full_tensor(), output_tensor)
326+
327+ @SupportedDevices(['Ascend910B'])
328+ @skipIfUnsupportMultiNPU(2)
329+ @with_comms
330+ def test_torch_npu_npu_conv2d_bias_is_None_input_shard0(self):
331+ mesh = self.build_device_mesh()
332+ 
333+ input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
334+ weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
335+ 
336+ input_dtensor = distribute_tensor(input_tensor, mesh, [Shard(0)])
337+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
338+ 
339+ bias = None
340+ 
341+ stride = (1, 1)
342+ padding = (1, 1)
343+ dilation = (1, 1)
344+ groups = 1
345+ 
346+ output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
347+ output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, bias, stride, padding, dilation, groups)
348+ self.assertEqual(output_dtensor.full_tensor(), output_tensor)
349+ 
350+ @SupportedDevices(['Ascend910B'])
351+ @skipIfUnsupportMultiNPU(2)
352+ @with_comms
353+ def test_torch_npu_npu_conv2d_bias_is_None_weight_shard0(self):
354+ mesh = self.build_device_mesh()
355+ 
356+ input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
357+ weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
358+ 
359+ input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
360+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Shard(0)])
361+ 
362+ bias = None
363+ 
364+ stride = (1, 1)
365+ padding = (1, 1)
366+ dilation = (1, 1)
367+ groups = 1
368+ 
369+ output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
370+ output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, bias, stride, padding, dilation, groups)
371+ self.assertEqual(output_dtensor.full_tensor(), output_tensor)
372+ 
373+ @SupportedDevices(['Ascend910B'])
374+ @skipIfUnsupportMultiNPU(2)
375+ @with_comms
376+ def test_torch_npu_npu_conv2d_backward_replicate(self):
377+ mesh = self.build_device_mesh()
378+ 
379+ input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
380+ weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
381+ 
382+ input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
383+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
384+ 
385+ bias = torch.randn(4, device="npu", requires_grad=True)
386+ d_bias = distribute_tensor(bias, mesh, [Replicate()])
387+ 
388+ stride = (1, 1)
389+ padding = (1, 1)
390+ dilation = (1, 1)
391+ groups = 1
392+ output_mask = [True, True, True]
393+ 
394+ output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
395+ grad_output = torch.ones_like(output_tensor, device="npu")
396+ grad_output_dtensor = distribute_tensor(grad_output, mesh, [Replicate()])
397+ 
398+ input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
399+ input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
400+ self.assertEqual(input_dgrad.full_tensor(), input_grad)
401+ self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
402+ self.assertEqual(bias_dgrad.full_tensor(), bias_grad)
403+ 
404+ @SupportedDevices(['Ascend910B'])
405+ @skipIfUnsupportMultiNPU(2)
406+ @with_comms
407+ def test_torch_npu_npu_conv2d_backward_bias_is_None_replicate(self):
408+ mesh = self.build_device_mesh()
409+ 
410+ input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
411+ weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
412+ 
413+ input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
414+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
415+ 
416+ bias = None
417+ 
418+ stride = (1, 1)
419+ padding = (1, 1)
420+ dilation = (1, 1)
421+ groups = 1
422+ output_mask = [True, True, False]
423+ 
424+ output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
425+ grad_output = torch.ones_like(output_tensor, device="npu")
426+ grad_output_dtensor = distribute_tensor(grad_output, mesh, [Replicate()])
427+ 
428+ input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
429+ input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
430+ self.assertEqual(input_dgrad.full_tensor(), input_grad)
431+ self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
432+
433+ @SupportedDevices(['Ascend910B'])
434+ @skipIfUnsupportMultiNPU(2)
435+ @with_comms
436+ def test_torch_npu_npu_conv2d_backward_input_shard0(self):
437+ mesh = self.build_device_mesh()
438+ 
439+ input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
440+ weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
441+ 
442+ input_dtensor = distribute_tensor(input_tensor, mesh, [Shard(0)])
443+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
444+ 
445+ bias = torch.randn(4, device="npu", requires_grad=True)
446+ 
447+ stride = (1, 1)
448+ padding = (1, 1)
449+ dilation = (1, 1)
450+ groups = 1
451+ output_mask = [True, True, True]
452+ 
453+ output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
454+ grad_output = torch.ones_like(output_tensor, device="npu")
455+ grad_output_dtensor = distribute_tensor(grad_output, mesh, [Shard(0)])
456+ 
457+ input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
458+ input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
459+ self.assertEqual(input_dgrad.full_tensor(), input_grad)
460+ self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
461+ self.assertEqual(bias_dgrad.full_tensor(), bias_grad)
462+
463+ @SupportedDevices(['Ascend910B'])
464+ @skipIfUnsupportMultiNPU(2)
465+ @with_comms
466+ def test_torch_npu_npu_conv2d_backward_bias_is_None_input_shard0(self):
467+ mesh = self.build_device_mesh()
468+ 
469+ input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
470+ weight_tensor = torch.randn(4, 3, 3, 3, device="npu", requires_grad=True)
471+ 
472+ input_dtensor = distribute_tensor(input_tensor, mesh, [Shard(0)])
473+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
474+ 
475+ bias = None
476+ 
477+ stride = (1, 1)
478+ padding = (1, 1)
479+ dilation = (1, 1)
480+ groups = 1
481+ output_mask = [True, True, False]
482+ 
483+ output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
484+ grad_output = torch.ones_like(output_tensor, device="npu")
485+ grad_output_dtensor = distribute_tensor(grad_output, mesh, [Shard(0)])
486+ 
487+ input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
488+ input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
489+ self.assertEqual(input_dgrad.full_tensor(), input_grad)
490+ self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
491+
492+ @SupportedDevices(['Ascend910B'])
493+ @skipIfUnsupportMultiNPU(2)
494+ @with_comms
495+ def test_torch_npu_npu_conv2d_backward_weight_shard0(self):
496+ mesh = self.build_device_mesh()
497+ 
498+ input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
499+ weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
500+ 
501+ input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
502+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Shard(0)])
503+ 
504+ bias = torch.randn(4, device="npu", requires_grad=True)
505+ 
506+ stride = (1, 1)
507+ padding = (1, 1)
508+ dilation = (1, 1)
509+ groups = 1
510+ output_mask = [True, True, True]
511+ 
512+ output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
513+ grad_output = torch.ones_like(output_tensor, device="npu")
514+ grad_output_dtensor = distribute_tensor(grad_output, mesh, [Shard(1)])
515+ 
516+ input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
517+ input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
518+ self.assertEqual(input_dgrad.full_tensor(), input_grad)
519+ self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
520+ self.assertEqual(bias_dgrad.full_tensor(), bias_grad)
521+ 
522+ @SupportedDevices(['Ascend910B'])
523+ @skipIfUnsupportMultiNPU(2)
524+ @with_comms
525+ def test_torch_npu_npu_conv2d_backward_bias_is_None_weight_shard0(self):
526+ mesh = self.build_device_mesh()
527+ 
528+ input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
529+ weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
530+ 
531+ input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
532+ weight_dtensor = distribute_tensor(weight_tensor, mesh, [Shard(0)])
533+ 
534+ bias = None
535+ stride = (1, 1)
536+ padding = (1, 1)
537+ dilation = (1, 1)
538+ groups = 1
539+ output_mask = [True, True, False]
540+ 
541+ output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
542+ grad_output = torch.ones_like(output_tensor, device="npu")
543+ grad_output_dtensor = distribute_tensor(grad_output, mesh, [Shard(1)])
544+ 
545+ input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
546+ input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
547+ self.assertEqual(input_dgrad.full_tensor(), input_grad)
548+ self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
549+ 
550+ 
551+class TestGroupedMatmulAdd(NPUDTensorTestBase):
552+ @SupportedDevices(['Ascend910B'])
553+ @skipIfUnsupportMultiNPU(2)
554+ @with_comms
555+ def test_torch_npu_npu_grouped_matmul_add__replicate(self):
556+ mesh = self.build_device_mesh()
557+ 
558+ x = torch.randn(8, 8, dtype=torch.float16, device="npu")
559+ weight = torch.randn(8, 4, dtype=torch.float16, device="npu")
560+ y = torch.randn(32, 4, dtype=torch.float, device="npu")
561+ group_list = torch.tensor([2, 4, 6, 8]).to(torch.int64).npu()
562+ x_dtensor = distribute_tensor(x, mesh, [Replicate()])
563+ weight_dtensor = distribute_tensor(weight, mesh, [Replicate()])
564+ y_dtensor = distribute_tensor(y, mesh, [Replicate()])
565+ group_list_dtensor = distribute_tensor(group_list, mesh, [Replicate()])
566+ transpose_x = True
567+ transpose_weight = False
568+ group_type = 2
569+ 
570+ torch_npu.npu_grouped_matmul_add_(y, x, weight, group_list, transpose_x=transpose_x, transpose_weight=transpose_weight, group_type=group_type)
571+ torch_npu.npu_grouped_matmul_add_(y_dtensor, x_dtensor, weight_dtensor, group_list_dtensor, transpose_x=transpose_x, transpose_weight=transpose_weight, group_type=group_type)
572+ self.assertEqual(y_dtensor.full_tensor(), y)
573+
574+ @SupportedDevices(['Ascend910B'])
575+ @skipIfUnsupportMultiNPU(2)
576+ @with_comms
577+ def test_torch_npu_npu_grouped_matmul_add__shard_D_weight(self):
578+ mesh = self.build_device_mesh()
579+ 
580+ x = torch.randn(8, 8, dtype=torch.float16, device="npu")
581+ weight = torch.randn(8, 4, dtype=torch.float16, device="npu")
582+ y = torch.randn(32, 4, dtype=torch.float, device="npu")
583+ group_list = torch.tensor([2, 4, 6, 8]).to(torch.int64).npu()
584+ x_dtensor = distribute_tensor(x, mesh, [Shard(1)])
585+ weight_dtensor = distribute_tensor(weight, mesh, [Shard(1)])
586+ y_dtensor = distribute_tensor(y, mesh, [Shard(1)])
587+ group_list_dtensor = distribute_tensor(group_list, mesh, [Replicate()])
588+ transpose_x = True
589+ transpose_weight = False
590+ group_type = 2
591+ 
592+ torch_npu.npu_grouped_matmul_add_(y, x, weight, group_list, transpose_x=transpose_x, transpose_weight=transpose_weight, group_type=group_type)
593+ torch_npu.npu_grouped_matmul_add_(y_dtensor, x_dtensor, weight_dtensor, group_list_dtensor, transpose_x=transpose_x, transpose_weight=transpose_weight, group_type=group_type)
594+ self.assertEqual(y_dtensor.full_tensor(), y)
595+ 
596+ 
597+class TestCrossEntropyLoss(NPUDTensorTestBase):
598+ def generate_data_cross_entropy_loss(self, N, C, input_strategy, target_strategy, weight_strategy=None):
599+ mesh = self.build_device_mesh()
600+ 
601+ x = torch.randn(N, C, device="npu", requires_grad=True)
602+ target = torch.arange(0, N, device="npu")
603+ input_dtensor = distribute_tensor(x, mesh, input_strategy)
604+ target_dtensor = distribute_tensor(target, mesh, target_strategy)
605+ 
606+ if weight_strategy:
607+ weight = torch.rand(C, device="npu")
608+ weight_dtensor = distribute_tensor(weight, mesh, weight_strategy)
609+ 
610+ input_tuple = (x, target, weight, input_dtensor, target_dtensor, weight_dtensor, mesh)
611+ 
612+ return input_tuple
613+ else:
614+ input_tuple = (x, target, input_dtensor, target_dtensor, mesh)
615+
616+ return input_tuple
617+ 
618+ 
619+ @SupportedDevices(['Ascend910B'])
620+ @skipIfUnsupportMultiNPU(2)
621+ @with_comms
622+ def test_torch_npu_npu_cross_entropy_loss_replicate(self):
623+ x, target, input_dtensor, target_dtensor, _ = self.generate_data_cross_entropy_loss(8, 8, [Replicate()], [Replicate()])
624+ 
625+ loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction="none")
626+ loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction="none")
627+ 
628+ self.assertEqual(loss_dtensor.full_tensor(), loss)
629+ self.assertEqual(log_prob_dtensor.full_tensor(), log_prob)
630+ 
631+
632+ @SupportedDevices(['Ascend910B'])
633+ @skipIfUnsupportMultiNPU(2)
634+ @with_comms
635+ def test_torch_npu_npu_cross_entropy_loss_input_shard0_not_evenly_shardable(self):
636+ x, target, input_dtensor, target_dtensor, _ = self.generate_data_cross_entropy_loss(7, 8, [Shard(0)], [Shard(0)])
637+ 
638+ loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction="mean")
639+ loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction="mean")
640+ 
641+ self.assertEqual(loss_dtensor.full_tensor(), loss)
642+ self.assertEqual(log_prob_dtensor.full_tensor(), log_prob)
643+ 
644+
645+ @SupportedDevices(['Ascend910B'])
646+ @skipIfUnsupportMultiNPU(2)
647+ @with_comms
648+ def test_torch_npu_npu_cross_entropy_loss_input_shard0_evenly_shardable(self):
649+ x, target, input_dtensor, target_dtensor, _ = self.generate_data_cross_entropy_loss(8, 8, [Shard(0)], [Shard(0)])
650+ 
651+ loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction="mean")
652+ loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction="mean")
653+ 
654+ self.assertEqual(loss_dtensor.full_tensor(), loss)
655+ self.assertEqual(log_prob_dtensor.full_tensor(), log_prob)
656+ 
657+ 
658+ @SupportedDevices(['Ascend910B'])
659+ @skipIfUnsupportMultiNPU(2)
660+ @with_comms
661+ def test_torch_npu_npu_cross_entropy_loss_input_shard0_evenly_shardable_weight(self):
662+ reductions = ["none", "sum"]
663+ x, target, weight, input_dtensor, target_dtensor, weight_dtensor, _ = self.generate_data_cross_entropy_loss(8, 8, [Shard(0)], [Shard(0)], [Replicate()])
664+ 
665+ for re in reductions:
666+ loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, weight, re)
667+ loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, weight_dtensor, re)
668+ 
669+ self.assertEqual(loss_dtensor.full_tensor(), loss)
670+ self.assertEqual(log_prob_dtensor.full_tensor(), log_prob)
671+ 
672+
673+ @SupportedDevices(['Ascend910B'])
674+ @skipIfUnsupportMultiNPU(2)
675+ @with_comms
676+ def test_torch_npu_npu_cross_entropy_loss_backward_replicate_reduction_is_mean(self):
677+ x, target, input_dtensor, target_dtensor, _ = self.generate_data_cross_entropy_loss(8, 8, [Replicate()], [Replicate()])
678+ 
679+ loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction="mean")
680+ loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction="mean")
681+ 
682+ loss.backward()
683+ loss_dtensor.backward()
684+ self.assertEqual(input_dtensor.grad.full_tensor(), x.grad)
685+ 
686+
687+ @SupportedDevices(['Ascend910B'])
688+ @skipIfUnsupportMultiNPU(2)
689+ @with_comms
690+ def test_torch_npu_npu_cross_entropy_loss_backward_input_shard0_reduction_is_none(self):
691+ reductions = ["none", "sum", "mean"]
692+ x, target, input_dtensor, target_dtensor, mesh = self.generate_data_cross_entropy_loss(8, 8, [Shard(0)], [Shard(0)])
693+
694+ for re in reductions:
695+ loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction=re)
696+ loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction=re)
697+ if re == "none":
698+ grad = torch.randn(loss.size(), device="npu")
699+ grad_dtensor = distribute_tensor(grad, mesh, [Shard(0)])
700+ 
701+ loss.backward(grad)
702+ loss_dtensor.backward(grad_dtensor)
703+ self.assertEqual(input_dtensor.grad.full_tensor(), x.grad)
704+ else:
705+ loss.backward()
706+ loss_dtensor.backward()
707+ self.assertEqual(input_dtensor.grad.full_tensor(), x.grad)
708+
709+ @SupportedDevices(['Ascend910B'])
710+ @skipIfUnsupportMultiNPU(2)
711+ @with_comms
712+ def test_torch_npu_npu_cross_entropy_loss_backward_input_shard1_reduction_is_sum(self):
713+ x, target, input_dtensor, target_dtensor, _ = self.generate_data_cross_entropy_loss(8, 8, [Shard(1)], [Shard(0)])
714+ 
715+ loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction="sum")
716+ loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction="sum")
717+
718+ loss.backward()
719+ loss_dtensor.backward()
720+ self.assertEqual(input_dtensor.grad.full_tensor(), x.grad)
721+ 
722+ 
723+class TestRepeatInterleaveSelfInt(NPUDTensorTestBase):
724+ def generate_data_repeat_interleave_self_int(self, size, repeats_value, input_strategy):
725+ mesh = self.build_device_mesh()
726+ 
727+ input_tensor = torch.randn(size, device="npu", requires_grad=True)
728+ input_dtensor = distribute_tensor(input_tensor, mesh, input_strategy)
729+ 
730+ result = (input_tensor, repeats_value, input_dtensor, mesh)
731+ 
732+ return result
733+ 
734+ @SupportedDevices(['Ascend910B'])
735+ @skipIfUnsupportMultiNPU(2)
736+ @with_comms
737+ def test_torch_repeat_interleave_self_int_replicate(self):
738+ input_tensor, repeats_value, input_dtensor, _ = self.generate_data_repeat_interleave_self_int((5, 5), 3, [Replicate()])
739+ 
740+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
741+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
742+
743+ self.assertEqual(output_dtensor.full_tensor(), output)
744+ 
745+ @SupportedDevices(['Ascend910B'])
746+ @skipIfUnsupportMultiNPU(2)
747+ @with_comms
748+ def test_torch_repeat_interleave_self_int_shard1(self):
749+ input_tensor, repeats_value, input_dtensor, _ = self.generate_data_repeat_interleave_self_int((8, 8), 3, [Shard(1)])
750+ 
751+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
752+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
753+
754+ self.assertEqual(output_dtensor.full_tensor(), output)
755+ 
756+ @SupportedDevices(['Ascend910B'])
757+ @skipIfUnsupportMultiNPU(2)
758+ @with_comms
759+ def test_torch_repeat_interleave_self_int_shard0(self):
760+ input_tensor, repeats_value, input_dtensor, _ = self.generate_data_repeat_interleave_self_int((8, 8), 3, [Shard(0)])
761+ 
762+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
763+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
764+
765+ self.assertEqual(output_dtensor.full_tensor(), output)
766+ 
767+ @SupportedDevices(['Ascend910B'])
768+ @skipIfUnsupportMultiNPU(2)
769+ @with_comms
770+ def test_torch_repeat_interleave_self_int_dim_is_None_shard0_is_evenly_shardable(self):
771+ input_tensor, repeats_value, input_dtensor, _ = self.generate_data_repeat_interleave_self_int((8, 5), 3, [Shard(0)])
772+ 
773+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value)
774+ output = torch.repeat_interleave(input_tensor, repeats_value)
775+
776+ self.assertEqual(output_dtensor.full_tensor(), output)
777+ 
778+ @SupportedDevices(['Ascend910B'])
779+ @skipIfUnsupportMultiNPU(2)
780+ @with_comms
781+ def test_torch_repeat_interleave_self_int_shard0_dim1_is_not_evenly_shardable(self):
782+ input_tensor, repeats_value, input_dtensor, _ = self.generate_data_repeat_interleave_self_int((5, 5), 3, [Shard(0)])
783+ 
784+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
785+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
786+
787+ self.assertEqual(output_dtensor.full_tensor(), output)
788+ 
789+ @SupportedDevices(['Ascend910B'])
790+ @skipIfUnsupportMultiNPU(2)
791+ @with_comms
792+ def test_torch_repeat_interleave_self_int_shard1_dim1_is_not_evenly_shardable(self):
793+ input_tensor, repeats_value, input_dtensor, _ = self.generate_data_repeat_interleave_self_int((5, 8), 3, [Shard(1)])
794+ 
795+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
796+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
797+
798+ self.assertEqual(output_dtensor.full_tensor(), output)
799+ 
800+ @SupportedDevices(['Ascend910B'])
801+ @skipIfUnsupportMultiNPU(2)
802+ @with_comms
803+ def test_torch_repeat_interleave_backward_self_int_replicate_dim1(self):
804+ sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
805+ 
806+ for size in sizes:
807+ input_tensor, repeats_value, input_dtensor, mesh = self.generate_data_repeat_interleave_self_int(size, 3, [Replicate()])
808+ 
809+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
810+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
811+ 
812+ grad_tensor = torch.randn(output.size(), device="npu")
813+ grad_dtensor = distribute_tensor(grad_tensor, mesh, [Replicate()])
814+ 
815+ output_dtensor.backward(grad_dtensor)
816+ output.backward(grad_tensor)
817+ self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
818+ 
819+ @SupportedDevices(['Ascend910B'])
820+ @skipIfUnsupportMultiNPU(2)
821+ @with_comms
822+ def test_torch_repeat_interleave_backward_self_int_replicate_shard0_dim1(self):
823+ sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
824+ 
825+ for size in sizes:
826+ input_tensor, repeats_value, input_dtensor, mesh = self.generate_data_repeat_interleave_self_int(size, 3, [Replicate()])
827+ 
828+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
829+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
830+ 
831+ grad_tensor = torch.randn(output.size(), device="npu")
832+ grad_dtensor = distribute_tensor(grad_tensor, mesh, [Shard(0)])
833+ 
834+ output_dtensor.backward(grad_dtensor)
835+ output.backward(grad_tensor)
836+ self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
837+ 
838+ @SupportedDevices(['Ascend910B'])
839+ @skipIfUnsupportMultiNPU(2)
840+ @with_comms
841+ def test_torch_repeat_interleave_backward_self_int_shard1_replicate_dim1(self):
842+ sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
843+ 
844+ for size in sizes:
845+ input_tensor, repeats_value, input_dtensor, mesh = self.generate_data_repeat_interleave_self_int(size, 3, [Shard(1)])
846+ 
847+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
848+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
849+ 
850+ grad_tensor = torch.randn(output.size(), device="npu")
851+ grad_dtensor = distribute_tensor(grad_tensor, mesh, [Replicate()])
852+ 
853+ output_dtensor.backward(grad_dtensor)
854+ output.backward(grad_tensor)
855+ self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
856+ 
857+ @SupportedDevices(['Ascend910B'])
858+ @skipIfUnsupportMultiNPU(2)
859+ @with_comms
860+ def test_torch_repeat_interleave_backward_self_int_replicate_dim_None(self):
861+ sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
862+ 
863+ for size in sizes:
864+ input_tensor, repeats_value, input_dtensor, mesh = self.generate_data_repeat_interleave_self_int(size, 3, [Replicate()])
865+ 
866+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value)
867+ output = torch.repeat_interleave(input_tensor, repeats_value)
868+ 
869+ grad_tensor = torch.randn(output.size(), device="npu")
870+ grad_dtensor = distribute_tensor(grad_tensor, mesh, [Replicate()])
871+ 
872+ output_dtensor.backward(grad_dtensor)
873+ output.backward(grad_tensor)
874+ self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
875+ 
876+ @SupportedDevices(['Ascend910B'])
877+ @skipIfUnsupportMultiNPU(2)
878+ @with_comms
879+ def test_torch_repeat_interleave_backward_self_int_shard00_dim1(self):
880+ sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
881+ 
882+ for size in sizes:
883+ input_tensor, repeats_value, input_dtensor, mesh = self.generate_data_repeat_interleave_self_int(size, 3, [Shard(0)])
884+ 
885+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
886+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
887+ 
888+ grad_tensor = torch.randn(output.size(), device="npu")
889+ grad_dtensor = distribute_tensor(grad_tensor, mesh, [Shard(0)])
890+ 
891+ output_dtensor.backward(grad_dtensor)
892+ output.backward(grad_tensor)
893+ self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
894+ 
895+ @SupportedDevices(['Ascend910B'])
896+ @skipIfUnsupportMultiNPU(2)
897+ @with_comms
898+ def test_torch_repeat_interleave_backward_self_int_shard01_dim1(self):
899+ sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
900+ 
901+ for size in sizes:
902+ input_tensor, repeats_value, input_dtensor, mesh = self.generate_data_repeat_interleave_self_int(size, 3, [Shard(0)])
903+ 
904+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
905+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
906+ 
907+ grad_tensor = torch.randn(output.size(), device="npu")
908+ grad_dtensor = distribute_tensor(grad_tensor, mesh, [Shard(1)])
909+ 
910+ output_dtensor.backward(grad_dtensor)
911+ output.backward(grad_tensor)
912+ self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
913+ 
914+ @SupportedDevices(['Ascend910B'])
915+ @skipIfUnsupportMultiNPU(2)
916+ @with_comms
917+ def test_torch_repeat_interleave_backward_self_int_shard10_dim1(self):
918+ sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
919+ 
920+ for size in sizes:
921+ input_tensor, repeats_value, input_dtensor, mesh = self.generate_data_repeat_interleave_self_int(size, 3, [Shard(1)])
922+ 
923+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
924+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
925+ 
926+ grad_tensor = torch.randn(output.size(), device="npu")
927+ grad_dtensor = distribute_tensor(grad_tensor, mesh, [Shard(0)])
928+ 
929+ output_dtensor.backward(grad_dtensor)
930+ output.backward(grad_tensor)
931+ self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
932+ 
933+ @SupportedDevices(['Ascend910B'])
934+ @skipIfUnsupportMultiNPU(2)
935+ @with_comms
936+ def test_torch_repeat_interleave_backward_self_int_shard11_dim1(self):
937+ sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
938+ 
939+ for size in sizes:
940+ input_tensor, repeats_value, input_dtensor, mesh = self.generate_data_repeat_interleave_self_int(size, 3, [Shard(1)])
941+ 
942+ output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
943+ output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
944+ 
945+ grad_tensor = torch.randn(output.size(), device="npu")
946+ grad_dtensor = distribute_tensor(grad_tensor, mesh, [Shard(1)])
947+ 
948+ output_dtensor.backward(grad_dtensor)
949+ output.backward(grad_tensor)
950+ self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
951+ 
952+ 
226instantiate_parametrized_tests(TestMathOps)953instantiate_parametrized_tests(TestMathOps)
227 954 
228 955 
@@ -1,4 +1,5 @@
1import itertools1import itertools
2+import numpy as np
2 3 
3import torch4import torch
4from torch.distributed._tensor import distribute_tensor, Replicate, Shard5from torch.distributed._tensor import distribute_tensor, Replicate, Shard
@@ -481,6 +482,338 @@ class TestGroupedMatMulOp(NPUDTensorTestBase):
481 test_placement_comb([comb[0]], [comb[1]], [comb[2]], [comb[2]])482 test_placement_comb([comb[0]], [comb[1]], [comb[2]], [comb[2]])
482 483 
483 484 
485+class TestApplyAdamW(NPUDTensorTestBase):
486+ @SupportedDevices(['Ascend910B'])
487+ @skipIfUnsupportMultiNPU(2)
488+ @with_comms
489+ def test_torch_npu_npu_apply_adam_w_replicate(self):
490+ mesh = self.build_device_mesh()
491+ 
492+ amsgrad = False
493+ maximize = True
494+ scalar_shape = [1]
495+ input_size = (21130, 512)
496+ 
497+ var_npu = torch.randn(input_size, device="npu")
498+ m_npu = torch.randn(input_size, device="npu")
499+ v_npu = torch.randn(input_size, device="npu")
500+ grad_npu = torch.randn(input_size, device="npu")
501+ 
502+ var_npu_dtensor = distribute_tensor(var_npu, mesh, [Replicate()])
503+ m_npu_dtensor = distribute_tensor(m_npu, mesh, [Replicate()])
504+ v_npu_dtensor = distribute_tensor(v_npu, mesh, [Replicate()])
505+ grad_npu_dtensor = distribute_tensor(grad_npu, mesh, [Replicate()])
506+
507+ np.random.seed(42)
508+ 
509+ beta1_power = np.random.uniform(0.0, 1.0, scalar_shape)
510+ beta2_power = np.random.uniform(0.0, 1.0, scalar_shape)
511+ lr = np.random.uniform(0.0001, 0.1, scalar_shape)
512+ weight_decay = np.random.uniform(0.001, 0.1, scalar_shape)
513+ beta1 = np.random.uniform(0.5, 1.0, scalar_shape)
514+ beta2 = np.random.uniform(0.5, 1.0, scalar_shape)
515+ eps = np.random.uniform(0.00001, 0.01, scalar_shape)
516+ max_grad_norm = None
517+
518+ var_ret_npu, m_ret_npu, v_ret_npu = torch_npu.npu_apply_adam_w(
519+ beta1_power[0],
520+ beta2_power[0],
521+ lr[0],
522+ weight_decay[0],
523+ beta1[0],
524+ beta2[0],
525+ eps[0],
526+ grad_npu,
527+ max_grad_norm,
528+ amsgrad,
529+ maximize,
530+ out=(var_npu, m_npu, v_npu),
531+ )
532+ 
533+ var_ret_npu_dtensor, m_ret_npu_dtensor, v_ret_npu_dtensor = torch_npu.npu_apply_adam_w(
534+ beta1_power[0],
535+ beta2_power[0],
536+ lr[0],
537+ weight_decay[0],
538+ beta1[0],
539+ beta2[0],
540+ eps[0],
541+ grad_npu_dtensor,
542+ max_grad_norm,
543+ amsgrad,
544+ maximize,
545+ out=(var_npu_dtensor, m_npu_dtensor, v_npu_dtensor),
546+ )
547+ 
548+ self.assertEqual(var_ret_npu_dtensor.full_tensor(), var_ret_npu)
549+ self.assertEqual(m_ret_npu_dtensor.full_tensor(), m_ret_npu)
550+ self.assertEqual(v_ret_npu_dtensor.full_tensor(), v_ret_npu)
551+ 
552+ @SupportedDevices(['Ascend910B'])
553+ @skipIfUnsupportMultiNPU(2)
554+ @with_comms
555+ def test_torch_npu_npu_apply_adam_w_shard00(self):
556+ mesh = self.build_device_mesh()
557+ 
558+ amsgrad = False
559+ maximize = True
560+ scalar_shape = [1]
561+ input_size = (21130, 512)
562+ 
563+ var_npu = torch.randn(input_size, device="npu")
564+ m_npu = torch.randn(input_size, device="npu")
565+ v_npu = torch.randn(input_size, device="npu")
566+ grad_npu = torch.randn(input_size, device="npu")
567+ 
568+ var_npu_dtensor = distribute_tensor(var_npu, mesh, [Shard(0)])
569+ m_npu_dtensor = distribute_tensor(m_npu, mesh, [Shard(0)])
570+ v_npu_dtensor = distribute_tensor(v_npu, mesh, [Shard(0)])
571+ grad_npu_dtensor = distribute_tensor(grad_npu, mesh, [Shard(0)])
572+
573+ np.random.seed(42)
574+ 
575+ beta1_power = np.random.uniform(0.0, 1.0, scalar_shape)
576+ beta2_power = np.random.uniform(0.0, 1.0, scalar_shape)
577+ lr = np.random.uniform(0.0001, 0.1, scalar_shape)
578+ weight_decay = np.random.uniform(0.001, 0.1, scalar_shape)
579+ beta1 = np.random.uniform(0.5, 1.0, scalar_shape)
580+ beta2 = np.random.uniform(0.5, 1.0, scalar_shape)
581+ eps = np.random.uniform(0.00001, 0.01, scalar_shape)
582+ max_grad_norm = None
583+
584+ var_ret_npu, m_ret_npu, v_ret_npu = torch_npu.npu_apply_adam_w(
585+ beta1_power[0],
586+ beta2_power[0],
587+ lr[0],
588+ weight_decay[0],
589+ beta1[0],
590+ beta2[0],
591+ eps[0],
592+ grad_npu,
593+ max_grad_norm,
594+ amsgrad,
595+ maximize,
596+ out=(var_npu, m_npu, v_npu),
597+ )
598+ 
599+ var_ret_npu_dtensor, m_ret_npu_dtensor, v_ret_npu_dtensor = torch_npu.npu_apply_adam_w(
600+ beta1_power[0],
601+ beta2_power[0],
602+ lr[0],
603+ weight_decay[0],
604+ beta1[0],
605+ beta2[0],
606+ eps[0],
607+ grad_npu_dtensor,
608+ max_grad_norm,
609+ amsgrad,
610+ maximize,
611+ out=(var_npu_dtensor, m_npu_dtensor, v_npu_dtensor),
612+ )
613+ 
614+ self.assertEqual(var_ret_npu_dtensor.full_tensor(), var_ret_npu)
615+ self.assertEqual(m_ret_npu_dtensor.full_tensor(), m_ret_npu)
616+ self.assertEqual(v_ret_npu_dtensor.full_tensor(), v_ret_npu)
617+ 
618+ @SupportedDevices(['Ascend910B'])
619+ @skipIfUnsupportMultiNPU(2)
620+ @with_comms
621+ def test_torch_npu_npu_apply_adam_w_shard01(self):
622+ mesh = self.build_device_mesh()
623+ 
624+ amsgrad = False
625+ maximize = True
626+ scalar_shape = [1]
627+ input_size = (21130, 512)
628+ 
629+ var_npu = torch.randn(input_size, device="npu")
630+ m_npu = torch.randn(input_size, device="npu")
631+ v_npu = torch.randn(input_size, device="npu")
632+ grad_npu = torch.randn(input_size, device="npu")
633+ 
634+ var_npu_dtensor = distribute_tensor(var_npu, mesh, [Shard(0)])
635+ m_npu_dtensor = distribute_tensor(m_npu, mesh, [Shard(0)])
636+ v_npu_dtensor = distribute_tensor(v_npu, mesh, [Shard(0)])
637+ grad_npu_dtensor = distribute_tensor(grad_npu, mesh, [Shard(1)])
638+
639+ np.random.seed(42)
640+ 
641+ beta1_power = np.random.uniform(0.0, 1.0, scalar_shape)
642+ beta2_power = np.random.uniform(0.0, 1.0, scalar_shape)
643+ lr = np.random.uniform(0.0001, 0.1, scalar_shape)
644+ weight_decay = np.random.uniform(0.001, 0.1, scalar_shape)
645+ beta1 = np.random.uniform(0.5, 1.0, scalar_shape)
646+ beta2 = np.random.uniform(0.5, 1.0, scalar_shape)
647+ eps = np.random.uniform(0.00001, 0.01, scalar_shape)
648+ max_grad_norm = None
649+
650+ var_ret_npu, m_ret_npu, v_ret_npu = torch_npu.npu_apply_adam_w(
651+ beta1_power[0],
652+ beta2_power[0],
653+ lr[0],
654+ weight_decay[0],
655+ beta1[0],
656+ beta2[0],
657+ eps[0],
658+ grad_npu,
659+ max_grad_norm,
660+ amsgrad,
661+ maximize,
662+ out=(var_npu, m_npu, v_npu),
663+ )
664+ 
665+ var_ret_npu_dtensor, m_ret_npu_dtensor, v_ret_npu_dtensor = torch_npu.npu_apply_adam_w(
666+ beta1_power[0],
667+ beta2_power[0],
668+ lr[0],
669+ weight_decay[0],
670+ beta1[0],
671+ beta2[0],
672+ eps[0],
673+ grad_npu_dtensor,
674+ max_grad_norm,
675+ amsgrad,
676+ maximize,
677+ out=(var_npu_dtensor, m_npu_dtensor, v_npu_dtensor),
678+ )
679+ 
680+ self.assertEqual(var_ret_npu_dtensor.full_tensor(), var_ret_npu)
681+ self.assertEqual(m_ret_npu_dtensor.full_tensor(), m_ret_npu)
682+ self.assertEqual(v_ret_npu_dtensor.full_tensor(), v_ret_npu)
683+ 
684+ @SupportedDevices(['Ascend910B'])
685+ @skipIfUnsupportMultiNPU(2)
686+ @with_comms
687+ def test_torch_npu_npu_apply_adam_w_shard10(self):
688+ mesh = self.build_device_mesh()
689+ 
690+ amsgrad = False
691+ maximize = True
692+ scalar_shape = [1]
693+ input_size = (21130, 512)
694+ 
695+ var_npu = torch.randn(input_size, device="npu")
696+ m_npu = torch.randn(input_size, device="npu")
697+ v_npu = torch.randn(input_size, device="npu")
698+ grad_npu = torch.randn(input_size, device="npu")
699+ 
700+ var_npu_dtensor = distribute_tensor(var_npu, mesh, [Shard(1)])
701+ m_npu_dtensor = distribute_tensor(m_npu, mesh, [Shard(1)])
702+ v_npu_dtensor = distribute_tensor(v_npu, mesh, [Shard(1)])
703+ grad_npu_dtensor = distribute_tensor(grad_npu, mesh, [Shard(0)])
704+
705+ np.random.seed(42)
706+ 
707+ beta1_power = np.random.uniform(0.0, 1.0, scalar_shape)
708+ beta2_power = np.random.uniform(0.0, 1.0, scalar_shape)
709+ lr = np.random.uniform(0.0001, 0.1, scalar_shape)
710+ weight_decay = np.random.uniform(0.001, 0.1, scalar_shape)
711+ beta1 = np.random.uniform(0.5, 1.0, scalar_shape)
712+ beta2 = np.random.uniform(0.5, 1.0, scalar_shape)
713+ eps = np.random.uniform(0.00001, 0.01, scalar_shape)
714+ max_grad_norm = None
715+
716+ var_ret_npu, m_ret_npu, v_ret_npu = torch_npu.npu_apply_adam_w(
717+ beta1_power[0],
718+ beta2_power[0],
719+ lr[0],
720+ weight_decay[0],
721+ beta1[0],
722+ beta2[0],
723+ eps[0],
724+ grad_npu,
725+ max_grad_norm,
726+ amsgrad,
727+ maximize,
728+ out=(var_npu, m_npu, v_npu),
729+ )
730+ 
731+ var_ret_npu_dtensor, m_ret_npu_dtensor, v_ret_npu_dtensor = torch_npu.npu_apply_adam_w(
732+ beta1_power[0],
733+ beta2_power[0],
734+ lr[0],
735+ weight_decay[0],
736+ beta1[0],
737+ beta2[0],
738+ eps[0],
739+ grad_npu_dtensor,
740+ max_grad_norm,
741+ amsgrad,
742+ maximize,
743+ out=(var_npu_dtensor, m_npu_dtensor, v_npu_dtensor),
744+ )
745+ 
746+ self.assertEqual(var_ret_npu_dtensor.full_tensor(), var_ret_npu)
747+ self.assertEqual(m_ret_npu_dtensor.full_tensor(), m_ret_npu)
748+ self.assertEqual(v_ret_npu_dtensor.full_tensor(), v_ret_npu)
749+ 
750+ @SupportedDevices(['Ascend910B'])
751+ @skipIfUnsupportMultiNPU(2)
752+ @with_comms
753+ def test_torch_npu_npu_apply_adam_w_shard11(self):
754+ mesh = self.build_device_mesh()
755+ 
756+ amsgrad = False
757+ maximize = True
758+ scalar_shape = [1]
759+ input_size = (21130, 512)
760+ 
761+ var_npu = torch.randn(input_size, device="npu")
762+ m_npu = torch.randn(input_size, device="npu")
763+ v_npu = torch.randn(input_size, device="npu")
764+ grad_npu = torch.randn(input_size, device="npu")
765+ 
766+ var_npu_dtensor = distribute_tensor(var_npu, mesh, [Shard(1)])
767+ m_npu_dtensor = distribute_tensor(m_npu, mesh, [Shard(1)])
768+ v_npu_dtensor = distribute_tensor(v_npu, mesh, [Shard(1)])
769+ grad_npu_dtensor = distribute_tensor(grad_npu, mesh, [Shard(1)])
770+
771+ np.random.seed(42)
772+ 
773+ beta1_power = np.random.uniform(0.0, 1.0, scalar_shape)
774+ beta2_power = np.random.uniform(0.0, 1.0, scalar_shape)
775+ lr = np.random.uniform(0.0001, 0.1, scalar_shape)
776+ weight_decay = np.random.uniform(0.001, 0.1, scalar_shape)
777+ beta1 = np.random.uniform(0.5, 1.0, scalar_shape)
778+ beta2 = np.random.uniform(0.5, 1.0, scalar_shape)
779+ eps = np.random.uniform(0.00001, 0.01, scalar_shape)
780+ max_grad_norm = None
781+
782+ var_ret_npu, m_ret_npu, v_ret_npu = torch_npu.npu_apply_adam_w(
783+ beta1_power[0],
784+ beta2_power[0],
785+ lr[0],
786+ weight_decay[0],
787+ beta1[0],
788+ beta2[0],
789+ eps[0],
790+ grad_npu,
791+ max_grad_norm,
792+ amsgrad,
793+ maximize,
794+ out=(var_npu, m_npu, v_npu),
795+ )
796+ 
797+ var_ret_npu_dtensor, m_ret_npu_dtensor, v_ret_npu_dtensor = torch_npu.npu_apply_adam_w(
798+ beta1_power[0],
799+ beta2_power[0],
800+ lr[0],
801+ weight_decay[0],
802+ beta1[0],
803+ beta2[0],
804+ eps[0],
805+ grad_npu_dtensor,
806+ max_grad_norm,
807+ amsgrad,
808+ maximize,
809+ out=(var_npu_dtensor, m_npu_dtensor, v_npu_dtensor),
810+ )
811+ 
812+ self.assertEqual(var_ret_npu_dtensor.full_tensor(), var_ret_npu)
813+ self.assertEqual(m_ret_npu_dtensor.full_tensor(), m_ret_npu)
814+ self.assertEqual(v_ret_npu_dtensor.full_tensor(), v_ret_npu)
815+ 
816+ 
484instantiate_parametrized_tests(TestGroupedMatMulOp)817instantiate_parametrized_tests(TestGroupedMatMulOp)
485 818 
486 819 
@@ -247,1058 +247,6 @@ class TestRegisterSharding(NPUDTensorTestBase):
247 else:247 else:
248 self.assertEqual(dist_dpse, dpse)248 self.assertEqual(dist_dpse, dpse)
249 249 
250- @skipIfUnsupportMultiNPU(4)
251- @with_comms
252- def test_torch_npu_npu_conv2d_replicate(self):
253- mesh = self.build_device_mesh()
254- 
255- input_tensor = torch.randn(3, 3, 224, 224, device="npu", requires_grad=True)
256- weight_tensor = torch.randn(64, 3, 3, 3, device="npu", requires_grad=True)
257- 
258- input_dtensor = distribute_tensor(input_tensor, mesh, [Shard(0)])
259- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
260- 
261- bias = torch.randn(64, device="npu", requires_grad=True)
262- d_bias = distribute_tensor(bias, mesh, [Replicate()])
263- 
264- 
265- stride = (1, 1)
266- padding = (1, 1)
267- dilation = (1, 1)
268- groups = 1
269- 
270- output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, d_bias, stride, padding, dilation, groups)
271- output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
272- 
273- self.assertEqual(output_dtensor.full_tensor(), output_tensor)
274- 
275- @skipIfUnsupportMultiNPU(4)
276- @with_comms
277- def test_torch_npu_npu_conv2d_weight_shard0(self):
278- mesh = self.build_device_mesh()
279- 
280- input_tensor = torch.randn(3, 3, 224, 224, device="npu", requires_grad=True)
281- weight_tensor = torch.randn(64, 3, 3, 3, device="npu", requires_grad=True)
282- 
283- input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
284- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Shard(0)])
285- 
286- bias = torch.randn(64, device="npu", requires_grad=True)
287- d_bias = distribute_tensor(bias, mesh, [Shard(0)])
288- 
289- 
290- stride = (1, 1)
291- padding = (1, 1)
292- dilation = (1, 1)
293- groups = 1
294- 
295- output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, d_bias, stride, padding, dilation, groups)
296- output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
297- 
298- self.assertEqual(output_dtensor.full_tensor(), output_tensor)
299- 
300- @skipIfUnsupportMultiNPU(4)
301- @with_comms
302- def test_torch_npu_npu_conv2d_input_shard1(self):
303- mesh = self.build_device_mesh()
304- 
305- input_tensor = torch.randn(8, 4, 224, 224, device="npu", requires_grad=True)
306- weight_tensor = torch.randn(64, 4, 3, 3, device="npu", requires_grad=True)
307- 
308- input_dtensor = distribute_tensor(input_tensor, mesh, [Shard(1)])
309- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Shard(1)])
310- 
311- bias = torch.randn(64, device="npu", requires_grad=True)
312- d_bias = distribute_tensor(bias, mesh, [Replicate()])
313- 
314- stride = (1, 1)
315- padding = (1, 1)
316- dilation = (1, 1)
317- groups = 1
318- 
319- output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, d_bias, stride, padding, dilation, groups)
320- output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
321- 
322- self.assertEqual(output_dtensor.full_tensor(), output_tensor)
323- 
324- @skipIfUnsupportMultiNPU(4)
325- @with_comms
326- def test_torch_npu_npu_conv2d_bias_is_None_replicate(self):
327- mesh = self.build_device_mesh()
328- 
329- input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
330- weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
331- 
332- input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
333- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
334- 
335- bias = None
336- 
337- stride = (1, 1)
338- padding = (1, 1)
339- dilation = (1, 1)
340- groups = 1
341- 
342- output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
343- output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, bias, stride, padding, dilation, groups)
344- self.assertEqual(output_dtensor.full_tensor(), output_tensor)
345-
346- @skipIfUnsupportMultiNPU(4)
347- @with_comms
348- def test_torch_npu_npu_conv2d_bias_is_None_input_shard0(self):
349- mesh = self.build_device_mesh()
350- 
351- input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
352- weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
353- 
354- input_dtensor = distribute_tensor(input_tensor, mesh, [Shard(0)])
355- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
356- 
357- bias = None
358- 
359- stride = (1, 1)
360- padding = (1, 1)
361- dilation = (1, 1)
362- groups = 1
363- 
364- output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
365- output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, bias, stride, padding, dilation, groups)
366- self.assertEqual(output_dtensor.full_tensor(), output_tensor)
367- 
368- @skipIfUnsupportMultiNPU(4)
369- @with_comms
370- def test_torch_npu_npu_conv2d_bias_is_None_weight_shard0(self):
371- mesh = self.build_device_mesh()
372- 
373- input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
374- weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
375- 
376- input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
377- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Shard(0)])
378- 
379- bias = None
380- 
381- stride = (1, 1)
382- padding = (1, 1)
383- dilation = (1, 1)
384- groups = 1
385- 
386- output_tensor = torch_npu.npu_conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
387- output_dtensor = torch_npu.npu_conv2d(input_dtensor, weight_dtensor, bias, stride, padding, dilation, groups)
388- self.assertEqual(output_dtensor.full_tensor(), output_tensor)
389- 
390- @skipIfUnsupportMultiNPU(4)
391- @with_comms
392- def test_torch_npu_npu_conv2d_backward_replicate(self):
393- mesh = self.build_device_mesh()
394- 
395- input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
396- weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
397- 
398- input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
399- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
400- 
401- bias = torch.randn(4, device="npu", requires_grad=True)
402- d_bias = distribute_tensor(bias, mesh, [Replicate()])
403- 
404- stride = (1, 1)
405- padding = (1, 1)
406- dilation = (1, 1)
407- groups = 1
408- output_mask = [True, True, True]
409- 
410- output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
411- grad_output = torch.ones_like(output_tensor, device="npu")
412- grad_output_dtensor = distribute_tensor(grad_output, mesh, [Replicate()])
413- 
414- input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
415- input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
416- self.assertEqual(input_dgrad.full_tensor(), input_grad)
417- self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
418- self.assertEqual(bias_dgrad.full_tensor(), bias_grad)
419- 
420- @skipIfUnsupportMultiNPU(4)
421- @with_comms
422- def test_torch_npu_npu_conv2d_backward_bias_is_None_replicate(self):
423- mesh = self.build_device_mesh()
424- 
425- input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
426- weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
427- 
428- input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
429- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
430- 
431- bias = None
432- 
433- stride = (1, 1)
434- padding = (1, 1)
435- dilation = (1, 1)
436- groups = 1
437- output_mask = [True, True, False]
438- 
439- output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
440- grad_output = torch.ones_like(output_tensor, device="npu")
441- grad_output_dtensor = distribute_tensor(grad_output, mesh, [Replicate()])
442- 
443- input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
444- input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
445- self.assertEqual(input_dgrad.full_tensor(), input_grad)
446- self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
447-
448- @skipIfUnsupportMultiNPU(4)
449- @with_comms
450- def test_torch_npu_npu_conv2d_backward_input_shard0(self):
451- mesh = self.build_device_mesh()
452- 
453- input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
454- weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
455- 
456- input_dtensor = distribute_tensor(input_tensor, mesh, [Shard(0)])
457- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
458- 
459- bias = torch.randn(4, device="npu", requires_grad=True)
460- 
461- stride = (1, 1)
462- padding = (1, 1)
463- dilation = (1, 1)
464- groups = 1
465- output_mask = [True, True, True]
466- 
467- output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
468- grad_output = torch.ones_like(output_tensor, device="npu")
469- grad_output_dtensor = distribute_tensor(grad_output, mesh, [Shard(0)])
470- 
471- input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
472- input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
473- self.assertEqual(input_dgrad.full_tensor(), input_grad)
474- self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
475- self.assertEqual(bias_dgrad.full_tensor(), bias_grad)
476-
477- @skipIfUnsupportMultiNPU(4)
478- @with_comms
479- def test_torch_npu_npu_conv2d_backward_bias_is_None_input_shard0(self):
480- mesh = self.build_device_mesh()
481- 
482- input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
483- weight_tensor = torch.randn(4, 3, 3, 3, device="npu", requires_grad=True)
484- 
485- input_dtensor = distribute_tensor(input_tensor, mesh, [Shard(0)])
486- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Replicate()])
487- 
488- bias = None
489- 
490- stride = (1, 1)
491- padding = (1, 1)
492- dilation = (1, 1)
493- groups = 1
494- output_mask = [True, True, False]
495- 
496- output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
497- grad_output = torch.ones_like(output_tensor, device="npu")
498- grad_output_dtensor = distribute_tensor(grad_output, mesh, [Shard(0)])
499- 
500- input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
501- input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
502- self.assertEqual(input_dgrad.full_tensor(), input_grad)
503- self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
504-
505- @skipIfUnsupportMultiNPU(4)
506- @with_comms
507- def test_torch_npu_npu_conv2d_backward_weight_shard0(self):
508- mesh = self.build_device_mesh()
509- 
510- input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
511- weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
512- 
513- input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
514- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Shard(0)])
515- 
516- bias = torch.randn(4, device="npu", requires_grad=True)
517- 
518- stride = (1, 1)
519- padding = (1, 1)
520- dilation = (1, 1)
521- groups = 1
522- output_mask = [True, True, True]
523- 
524- output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
525- grad_output = torch.ones_like(output_tensor, device="npu")
526- grad_output_dtensor = distribute_tensor(grad_output, mesh, [Shard(1)])
527- 
528- input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
529- input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
530- self.assertEqual(input_dgrad.full_tensor(), input_grad)
531- self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
532- self.assertEqual(bias_dgrad.full_tensor(), bias_grad)
533- 
534- @skipIfUnsupportMultiNPU(4)
535- @with_comms
536- def test_torch_npu_npu_conv2d_backward_bias_is_None_weight_shard0(self):
537- mesh = self.build_device_mesh()
538- 
539- input_tensor = torch.randn(4, 3, 28, 28, device="npu", requires_grad=True)
540- weight_tensor = torch.randn(4, 3, 4, 4, device="npu", requires_grad=True)
541- 
542- input_dtensor = distribute_tensor(input_tensor, mesh, [Replicate()])
543- weight_dtensor = distribute_tensor(weight_tensor, mesh, [Shard(0)])
544- 
545- bias = None
546- stride = (1, 1)
547- padding = (1, 1)
548- dilation = (1, 1)
549- groups = 1
550- output_mask = [True, True, False]
551- 
552- output_tensor = torch.nn.functional.conv2d(input_tensor, weight_tensor, bias, stride, padding, dilation, groups)
553- grad_output = torch.ones_like(output_tensor, device="npu")
554- grad_output_dtensor = distribute_tensor(grad_output, mesh, [Shard(1)])
555- 
556- input_grad, weight_grad, bias_grad = torch_npu.npu_conv2d_backward(input_tensor, grad_output, weight_tensor, stride, padding, dilation, groups, output_mask)
557- input_dgrad, weight_dgrad, bias_dgrad = torch_npu.npu_conv2d_backward(input_dtensor, grad_output_dtensor, weight_dtensor, stride, padding, dilation, groups, output_mask)
558- self.assertEqual(input_dgrad.full_tensor(), input_grad)
559- self.assertEqual(weight_dgrad.full_tensor(), weight_grad)
560-
561- @skipIfUnsupportMultiNPU(4)
562- @with_comms
563- def test_torch_npu_npu_grouped_matmul_add__replicate(self):
564- mesh = self.build_device_mesh()
565- 
566- x = torch.randn(8, 8, dtype=torch.float16, device="npu")
567- weight = torch.randn(8, 4, dtype=torch.float16, device="npu")
568- y = torch.randn(32, 4, dtype=torch.float, device="npu")
569- group_list = torch.tensor([2, 4, 6, 8]).to(torch.int64).npu()
570- x_dtensor = distribute_tensor(x, mesh, [Replicate()])
571- weight_dtensor = distribute_tensor(weight, mesh, [Replicate()])
572- y_dtensor = distribute_tensor(y, mesh, [Replicate()])
573- group_list_dtensor = distribute_tensor(group_list, mesh, [Replicate()])
574- transpose_x = True
575- transpose_weight = False
576- group_type = 1
577- 
578- torch_npu.npu_grouped_matmul_add_(y, x, weight, group_list, transpose_x=transpose_x, transpose_weight=transpose_weight, group_type=group_type)
579- torch_npu.npu_grouped_matmul_add_(y_dtensor, x_dtensor, weight_dtensor, group_list_dtensor, transpose_x=transpose_x, transpose_weight=transpose_weight, group_type=group_type)
580- self.assertEqual(y_dtensor.full_tensor(), y)
581-
582- @skipIfUnsupportMultiNPU(4)
583- @with_comms
584- def test_torch_npu_npu_grouped_matmul_add__shard_D_weight(self):
585- mesh = self.build_device_mesh()
586- 
587- x = torch.randn(8, 8, dtype=torch.float16, device="npu")
588- weight = torch.randn(8, 4, dtype=torch.float16, device="npu")
589- y = torch.randn(32, 4, dtype=torch.float, device="npu")
590- group_list = torch.tensor([2, 4, 6, 8]).to(torch.int64).npu()
591- x_dtensor = distribute_tensor(x, mesh, [Shard(1)])
592- weight_dtensor = distribute_tensor(weight, mesh, [Shard(1)])
593- y_dtensor = distribute_tensor(y, mesh, [Shard(1)])
594- group_list_dtensor = distribute_tensor(group_list, mesh, [Replicate()])
595- transpose_x = True
596- transpose_weight = False
597- group_type = 1
598- 
599- torch_npu.npu_grouped_matmul_add_(y, x, weight, group_list, transpose_x=transpose_x, transpose_weight=transpose_weight, group_type=group_type)
600- torch_npu.npu_grouped_matmul_add_(y_dtensor, x_dtensor, weight_dtensor, group_list_dtensor, transpose_x=transpose_x, transpose_weight=transpose_weight, group_type=group_type)
601- self.assertEqual(y_dtensor.full_tensor(), y)
602-
603- @skipIfUnsupportMultiNPU(4)
604- @with_comms
605- def test_torch_npu_npu_grouped_matmul_add__shard_D_x(self):
606- mesh = self.build_device_mesh()
607- 
608- x = torch.randn(8, 8, dtype=torch.float16, device="npu")
609- weight = torch.randn(8, 4, dtype=torch.float16, device="npu")
610- y = torch.randn(32, 4, dtype=torch.float, device="npu")
611- group_list = torch.tensor([2, 4, 6, 8]).to(torch.int64).npu()
612- x_dtensor = distribute_tensor(x, mesh, [Shard(0)])
613- weight_dtensor = distribute_tensor(weight, mesh, [Shard(1)])
614- y_dtensor = distribute_tensor(y, mesh, [Shard(0)])
615- group_list_dtensor = distribute_tensor(group_list, mesh, [Replicate()])
616- transpose_x = True
617- transpose_weight = False
618- group_type = 1
619- 
620- with self.assertRaises(RuntimeError) as cm:
621- torch_npu.npu_grouped_matmul_add_(y_dtensor, x_dtensor, weight_dtensor, group_list_dtensor, transpose_x=transpose_x, transpose_weight=transpose_weight, group_type=group_type)
622- 
623- err = cm.exception
624- self.assertIn("Sharding propagation failed for Op", str(err))
625- 
626- @skipIfUnsupportMultiNPU(4)
627- @with_comms
628- def test_torch_npu_npu_apply_adam_w_replicate(self):
629- mesh = self.build_device_mesh()
630- 
631- amsgrad = False
632- maximize = True
633- scalar_shape = [1]
634- input_size = (21130, 512)
635- 
636- var_npu = torch.randn(input_size, device="npu")
637- m_npu = torch.randn(input_size, device="npu")
638- v_npu = torch.randn(input_size, device="npu")
639- grad_npu = torch.randn(input_size, device="npu")
640- 
641- var_npu_dtensor = distribute_tensor(var_npu, mesh, [Replicate()])
642- m_npu_dtensor = distribute_tensor(m_npu, mesh, [Replicate()])
643- v_npu_dtensor = distribute_tensor(v_npu, mesh, [Replicate()])
644- grad_npu_dtensor = distribute_tensor(grad_npu, mesh, [Replicate()])
645-
646- np.random.seed(42)
647- 
648- beta1_power = np.random.uniform(0.0, 1.0, scalar_shape)
649- beta2_power = np.random.uniform(0.0, 1.0, scalar_shape)
650- lr = np.random.uniform(0.0001, 0.1, scalar_shape)
651- weight_decay = np.random.uniform(0.001, 0.1, scalar_shape)
652- beta1 = np.random.uniform(0.5, 1.0, scalar_shape)
653- beta2 = np.random.uniform(0.5, 1.0, scalar_shape)
654- eps = np.random.uniform(0.00001, 0.01, scalar_shape)
655- max_grad_norm = None
656-
657- var_ret_npu, m_ret_npu, v_ret_npu = torch_npu.npu_apply_adam_w(
658- beta1_power[0],
659- beta2_power[0],
660- lr[0],
661- weight_decay[0],
662- beta1[0],
663- beta2[0],
664- eps[0],
665- grad_npu,
666- max_grad_norm,
667- amsgrad,
668- maximize,
669- out=(var_npu, m_npu, v_npu),
670- )
671- 
672- var_ret_npu_dtensor, m_ret_npu_dtensor, v_ret_npu_dtensor = torch_npu.npu_apply_adam_w(
673- beta1_power[0],
674- beta2_power[0],
675- lr[0],
676- weight_decay[0],
677- beta1[0],
678- beta2[0],
679- eps[0],
680- grad_npu_dtensor,
681- max_grad_norm,
682- amsgrad,
683- maximize,
684- out=(var_npu_dtensor, m_npu_dtensor, v_npu_dtensor),
685- )
686- 
687- self.assertEqual(var_ret_npu_dtensor.full_tensor(), var_ret_npu)
688- self.assertEqual(m_ret_npu_dtensor.full_tensor(), m_ret_npu)
689- self.assertEqual(v_ret_npu_dtensor.full_tensor(), v_ret_npu)
690- 
691- @skipIfUnsupportMultiNPU(4)
692- @with_comms
693- def test_torch_npu_npu_apply_adam_w_shard00(self):
694- mesh = self.build_device_mesh()
695- 
696- amsgrad = False
697- maximize = True
698- scalar_shape = [1]
699- input_size = (21130, 512)
700- 
701- var_npu = torch.randn(input_size, device="npu")
702- m_npu = torch.randn(input_size, device="npu")
703- v_npu = torch.randn(input_size, device="npu")
704- grad_npu = torch.randn(input_size, device="npu")
705- 
706- var_npu_dtensor = distribute_tensor(var_npu, mesh, [Shard(0)])
707- m_npu_dtensor = distribute_tensor(m_npu, mesh, [Shard(0)])
708- v_npu_dtensor = distribute_tensor(v_npu, mesh, [Shard(0)])
709- grad_npu_dtensor = distribute_tensor(grad_npu, mesh, [Shard(0)])
710-
711- np.random.seed(42)
712- 
713- beta1_power = np.random.uniform(0.0, 1.0, scalar_shape)
714- beta2_power = np.random.uniform(0.0, 1.0, scalar_shape)
715- lr = np.random.uniform(0.0001, 0.1, scalar_shape)
716- weight_decay = np.random.uniform(0.001, 0.1, scalar_shape)
717- beta1 = np.random.uniform(0.5, 1.0, scalar_shape)
718- beta2 = np.random.uniform(0.5, 1.0, scalar_shape)
719- eps = np.random.uniform(0.00001, 0.01, scalar_shape)
720- max_grad_norm = None
721-
722- var_ret_npu, m_ret_npu, v_ret_npu = torch_npu.npu_apply_adam_w(
723- beta1_power[0],
724- beta2_power[0],
725- lr[0],
726- weight_decay[0],
727- beta1[0],
728- beta2[0],
729- eps[0],
730- grad_npu,
731- max_grad_norm,
732- amsgrad,
733- maximize,
734- out=(var_npu, m_npu, v_npu),
735- )
736- 
737- var_ret_npu_dtensor, m_ret_npu_dtensor, v_ret_npu_dtensor = torch_npu.npu_apply_adam_w(
738- beta1_power[0],
739- beta2_power[0],
740- lr[0],
741- weight_decay[0],
742- beta1[0],
743- beta2[0],
744- eps[0],
745- grad_npu_dtensor,
746- max_grad_norm,
747- amsgrad,
748- maximize,
749- out=(var_npu_dtensor, m_npu_dtensor, v_npu_dtensor),
750- )
751- 
752- self.assertEqual(var_ret_npu_dtensor.full_tensor(), var_ret_npu)
753- self.assertEqual(m_ret_npu_dtensor.full_tensor(), m_ret_npu)
754- self.assertEqual(v_ret_npu_dtensor.full_tensor(), v_ret_npu)
755- 
756- @skipIfUnsupportMultiNPU(4)
757- @with_comms
758- def test_torch_npu_npu_apply_adam_w_shard01(self):
759- mesh = self.build_device_mesh()
760- 
761- amsgrad = False
762- maximize = True
763- scalar_shape = [1]
764- input_size = (21130, 512)
765- 
766- var_npu = torch.randn(input_size, device="npu")
767- m_npu = torch.randn(input_size, device="npu")
768- v_npu = torch.randn(input_size, device="npu")
769- grad_npu = torch.randn(input_size, device="npu")
770- 
771- var_npu_dtensor = distribute_tensor(var_npu, mesh, [Shard(0)])
772- m_npu_dtensor = distribute_tensor(m_npu, mesh, [Shard(0)])
773- v_npu_dtensor = distribute_tensor(v_npu, mesh, [Shard(0)])
774- grad_npu_dtensor = distribute_tensor(grad_npu, mesh, [Shard(1)])
775-
776- np.random.seed(42)
777- 
778- beta1_power = np.random.uniform(0.0, 1.0, scalar_shape)
779- beta2_power = np.random.uniform(0.0, 1.0, scalar_shape)
780- lr = np.random.uniform(0.0001, 0.1, scalar_shape)
781- weight_decay = np.random.uniform(0.001, 0.1, scalar_shape)
782- beta1 = np.random.uniform(0.5, 1.0, scalar_shape)
783- beta2 = np.random.uniform(0.5, 1.0, scalar_shape)
784- eps = np.random.uniform(0.00001, 0.01, scalar_shape)
785- max_grad_norm = None
786-
787- var_ret_npu, m_ret_npu, v_ret_npu = torch_npu.npu_apply_adam_w(
788- beta1_power[0],
789- beta2_power[0],
790- lr[0],
791- weight_decay[0],
792- beta1[0],
793- beta2[0],
794- eps[0],
795- grad_npu,
796- max_grad_norm,
797- amsgrad,
798- maximize,
799- out=(var_npu, m_npu, v_npu),
800- )
801- 
802- var_ret_npu_dtensor, m_ret_npu_dtensor, v_ret_npu_dtensor = torch_npu.npu_apply_adam_w(
803- beta1_power[0],
804- beta2_power[0],
805- lr[0],
806- weight_decay[0],
807- beta1[0],
808- beta2[0],
809- eps[0],
810- grad_npu_dtensor,
811- max_grad_norm,
812- amsgrad,
813- maximize,
814- out=(var_npu_dtensor, m_npu_dtensor, v_npu_dtensor),
815- )
816- 
817- self.assertEqual(var_ret_npu_dtensor.full_tensor(), var_ret_npu)
818- self.assertEqual(m_ret_npu_dtensor.full_tensor(), m_ret_npu)
819- self.assertEqual(v_ret_npu_dtensor.full_tensor(), v_ret_npu)
820- 
821- @skipIfUnsupportMultiNPU(4)
822- @with_comms
823- def test_torch_npu_npu_apply_adam_w_shard10(self):
824- mesh = self.build_device_mesh()
825- 
826- amsgrad = False
827- maximize = True
828- scalar_shape = [1]
829- input_size = (21130, 512)
830- 
831- var_npu = torch.randn(input_size, device="npu")
832- m_npu = torch.randn(input_size, device="npu")
833- v_npu = torch.randn(input_size, device="npu")
834- grad_npu = torch.randn(input_size, device="npu")
835- 
836- var_npu_dtensor = distribute_tensor(var_npu, mesh, [Shard(1)])
837- m_npu_dtensor = distribute_tensor(m_npu, mesh, [Shard(1)])
838- v_npu_dtensor = distribute_tensor(v_npu, mesh, [Shard(1)])
839- grad_npu_dtensor = distribute_tensor(grad_npu, mesh, [Shard(0)])
840-
841- np.random.seed(42)
842- 
843- beta1_power = np.random.uniform(0.0, 1.0, scalar_shape)
844- beta2_power = np.random.uniform(0.0, 1.0, scalar_shape)
845- lr = np.random.uniform(0.0001, 0.1, scalar_shape)
846- weight_decay = np.random.uniform(0.001, 0.1, scalar_shape)
847- beta1 = np.random.uniform(0.5, 1.0, scalar_shape)
848- beta2 = np.random.uniform(0.5, 1.0, scalar_shape)
849- eps = np.random.uniform(0.00001, 0.01, scalar_shape)
850- max_grad_norm = None
851-
852- var_ret_npu, m_ret_npu, v_ret_npu = torch_npu.npu_apply_adam_w(
853- beta1_power[0],
854- beta2_power[0],
855- lr[0],
856- weight_decay[0],
857- beta1[0],
858- beta2[0],
859- eps[0],
860- grad_npu,
861- max_grad_norm,
862- amsgrad,
863- maximize,
864- out=(var_npu, m_npu, v_npu),
865- )
866- 
867- var_ret_npu_dtensor, m_ret_npu_dtensor, v_ret_npu_dtensor = torch_npu.npu_apply_adam_w(
868- beta1_power[0],
869- beta2_power[0],
870- lr[0],
871- weight_decay[0],
872- beta1[0],
873- beta2[0],
874- eps[0],
875- grad_npu_dtensor,
876- max_grad_norm,
877- amsgrad,
878- maximize,
879- out=(var_npu_dtensor, m_npu_dtensor, v_npu_dtensor),
880- )
881- 
882- self.assertEqual(var_ret_npu_dtensor.full_tensor(), var_ret_npu)
883- self.assertEqual(m_ret_npu_dtensor.full_tensor(), m_ret_npu)
884- self.assertEqual(v_ret_npu_dtensor.full_tensor(), v_ret_npu)
885- 
886- @skipIfUnsupportMultiNPU(4)
887- @with_comms
888- def test_torch_npu_npu_apply_adam_w_shard11(self):
889- mesh = self.build_device_mesh()
890- 
891- amsgrad = False
892- maximize = True
893- scalar_shape = [1]
894- input_size = (21130, 512)
895- 
896- var_npu = torch.randn(input_size, device="npu")
897- m_npu = torch.randn(input_size, device="npu")
898- v_npu = torch.randn(input_size, device="npu")
899- grad_npu = torch.randn(input_size, device="npu")
900- 
901- var_npu_dtensor = distribute_tensor(var_npu, mesh, [Shard(1)])
902- m_npu_dtensor = distribute_tensor(m_npu, mesh, [Shard(1)])
903- v_npu_dtensor = distribute_tensor(v_npu, mesh, [Shard(1)])
904- grad_npu_dtensor = distribute_tensor(grad_npu, mesh, [Shard(1)])
905-
906- np.random.seed(42)
907- 
908- beta1_power = np.random.uniform(0.0, 1.0, scalar_shape)
909- beta2_power = np.random.uniform(0.0, 1.0, scalar_shape)
910- lr = np.random.uniform(0.0001, 0.1, scalar_shape)
911- weight_decay = np.random.uniform(0.001, 0.1, scalar_shape)
912- beta1 = np.random.uniform(0.5, 1.0, scalar_shape)
913- beta2 = np.random.uniform(0.5, 1.0, scalar_shape)
914- eps = np.random.uniform(0.00001, 0.01, scalar_shape)
915- max_grad_norm = None
916-
917- var_ret_npu, m_ret_npu, v_ret_npu = torch_npu.npu_apply_adam_w(
918- beta1_power[0],
919- beta2_power[0],
920- lr[0],
921- weight_decay[0],
922- beta1[0],
923- beta2[0],
924- eps[0],
925- grad_npu,
926- max_grad_norm,
927- amsgrad,
928- maximize,
929- out=(var_npu, m_npu, v_npu),
930- )
931- 
932- var_ret_npu_dtensor, m_ret_npu_dtensor, v_ret_npu_dtensor = torch_npu.npu_apply_adam_w(
933- beta1_power[0],
934- beta2_power[0],
935- lr[0],
936- weight_decay[0],
937- beta1[0],
938- beta2[0],
939- eps[0],
940- grad_npu_dtensor,
941- max_grad_norm,
942- amsgrad,
943- maximize,
944- out=(var_npu_dtensor, m_npu_dtensor, v_npu_dtensor),
945- )
946- 
947- self.assertEqual(var_ret_npu_dtensor.full_tensor(), var_ret_npu)
948- self.assertEqual(m_ret_npu_dtensor.full_tensor(), m_ret_npu)
949- self.assertEqual(v_ret_npu_dtensor.full_tensor(), v_ret_npu)
950- 
951- @with_comms
952- def generate_data_cross_entropy_loss(self, N, C, input_strategy, target_strategy, weight_strategy=None):
953- mesh = self.build_device_mesh()
954- 
955- x = torch.randn(N, C, device="npu", requires_grad=True)
956- target = torch.arange(0, N, device="npu")
957- input_dtensor = distribute_tensor(x, mesh, input_strategy)
958- target_dtensor = distribute_tensor(target, mesh, target_strategy)
959- 
960- if weight_strategy:
961- weight = torch.rand(C, device="npu")
962- weight_dtensor = distribute_tensor(weight, mesh, weight_strategy)
963- 
964- input_tuple = (x, target, weight, input_dtensor, target_dtensor, weight_dtensor)
965- 
966- return input_tuple
967- else:
968- input_tuple = (x, target, input_dtensor, target_dtensor)
969-
970- return input_tuple
971- 
972- 
973- @skipIfUnsupportMultiNPU(4)
974- @with_comms
975- def test_torch_npu_npu_cross_entropy_loss_replicate(self):
976- x, target, input_dtensor, target_dtensor = self.generate_data_cross_entropy_loss(8, 8, [Replicate()], [Replicate()])
977- 
978- loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction="none")
979- loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction="none")
980- 
981- self.assertEqual(loss_dtensor.full_tensor(), loss)
982- self.assertEqual(log_prob_dtensor.full_tensor(), log_prob)
983- 
984-
985- @skipIfUnsupportMultiNPU(4)
986- @with_comms
987- def test_torch_npu_npu_cross_entropy_loss_input_shard0_not_evenly_shardable(self):
988- x, target, input_dtensor, target_dtensor = self.generate_data_cross_entropy_loss(7, 8, [Shard(0)], [Shard(0)])
989- 
990- loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction="mean")
991- loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction="mean")
992- 
993- self.assertEqual(loss_dtensor.full_tensor(), loss)
994- self.assertEqual(log_prob_dtensor.full_tensor(), log_prob)
995- 
996-
997- @skipIfUnsupportMultiNPU(4)
998- @with_comms
999- def test_torch_npu_npu_cross_entropy_loss_input_shard0_evenly_shardable(self):
1000- x, target, input_dtensor, target_dtensor = self.generate_data_cross_entropy_loss(8, 8, [Shard(0)], [Shard(0)])
1001- 
1002- loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction="mean")
1003- loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction="mean")
1004- 
1005- self.assertEqual(loss_dtensor.full_tensor(), loss)
1006- self.assertEqual(log_prob_dtensor.full_tensor(), log_prob)
1007- 
1008- 
1009- @skipIfUnsupportMultiNPU(4)
1010- @with_comms
1011- def test_torch_npu_npu_cross_entropy_loss_input_shard0_evenly_shardable_weight(self):
1012- reductions = ["none", "sum"]
1013- x, target, weight, input_dtensor, target_dtensor, weight_dtensor = self.generate_data_cross_entropy_loss(8, 8, [Shard(0)], [Shard(0)], [Replicate()])
1014- 
1015- for re in reductions:
1016- loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, weight, re)
1017- loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, weight_dtensor, re)
1018- 
1019- self.assertEqual(loss_dtensor.full_tensor(), loss)
1020- self.assertEqual(log_prob_dtensor.full_tensor(), log_prob)
1021- 
1022-
1023- @skipIfUnsupportMultiNPU(4)
1024- @with_comms
1025- def test_torch_npu_npu_cross_entropy_loss_backward_replicate_reduction_is_mean(self):
1026- x, target, input_dtensor, target_dtensor = self.generate_data_cross_entropy_loss(8, 8, [Replicate()], [Replicate()])
1027- 
1028- loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction="mean")
1029- loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction="mean")
1030- 
1031- loss.backward()
1032- loss_dtensor.backward()
1033- self.assertEqual(input_dtensor.grad.full_tensor(), x.grad)
1034- 
1035-
1036- @skipIfUnsupportMultiNPU(4)
1037- @with_comms
1038- def test_torch_npu_npu_cross_entropy_loss_backward_input_shard0_reduction_is_none(self):
1039- reductions = ["none", "sum", "mean"]
1040- x, target, input_dtensor, target_dtensor = self.generate_data_cross_entropy_loss(8, 8, [Shard(0)], [Shard(0)])
1041-
1042- for re in reductions:
1043- loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction=re)
1044- loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction=re)
1045- if re == "none":
1046- loss.backward()
1047- loss_dtensor.backward()
1048- self.assertEqual(input_dtensor.grad.full_tensor(), x.grad)
1049- else:
1050- grad = torch.randn(loss.size(), device="npu")
1051- grad_dtensor = distribute_tensor(grad, input_dtensor.mesh, [Shard(0)])
1052- 
1053- loss.backward(grad)
1054- loss_dtensor.backward(grad_dtensor)
1055- self.assertEqual(input_dtensor.grad.full_tensor(), x.grad)
1056-
1057- @skipIfUnsupportMultiNPU(4)
1058- @with_comms
1059- def test_torch_npu_npu_cross_entropy_loss_backward_input_shard1_reduction_is_sum(self):
1060- x, target, input_dtensor, target_dtensor = self.generate_data_cross_entropy_loss(8, 8, [Shard(1)], [Shard(0)])
1061- 
1062- loss, log_prob, _, _ = torch_npu.npu_cross_entropy_loss(x, target, reduction="sum")
1063- loss_dtensor, log_prob_dtensor, _, _ = torch_npu.npu_cross_entropy_loss(input_dtensor, target_dtensor, reduction="sum")
1064-
1065- loss.backward()
1066- loss_dtensor.backward()
1067- self.assertEqual(input_dtensor.grad.full_tensor(), x.grad)
1068-
1069- @with_comms
1070- def generate_data_repeat_interleave_self_int(self, size, repeats_value, input_strategy):
1071- mesh = self.build_device_mesh()
1072- 
1073- input_tensor = torch.randn(size, device="npu", requires_grad=True)
1074- input_dtensor = distribute_tensor(input_tensor, mesh, input_strategy)
1075- 
1076- return input_tensor, repeats_value, input_dtensor
1077- 
1078- @skipIfUnsupportMultiNPU(4)
1079- @with_comms
1080- def test_torch_repeat_interleave_self_int_replicate(self):
1081- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int((5, 5), 3, [Replicate()])
1082- 
1083- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1084- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1085-
1086- self.assertEqual(output_dtensor.full_tensor(), output)
1087- 
1088- @skipIfUnsupportMultiNPU(4)
1089- @with_comms
1090- def test_torch_repeat_interleave_self_int_shard1(self):
1091- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int((8, 8), 3, [Shard(1)])
1092- 
1093- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1094- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1095-
1096- self.assertEqual(output_dtensor.full_tensor(), output)
1097- 
1098- @skipIfUnsupportMultiNPU(4)
1099- @with_comms
1100- def test_torch_repeat_interleave_self_int_shard0(self):
1101- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int((8, 8), 3, [Shard(0)])
1102- 
1103- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1104- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1105-
1106- self.assertEqual(output_dtensor.full_tensor(), output)
1107- 
1108- @skipIfUnsupportMultiNPU(4)
1109- @with_comms
1110- def test_torch_repeat_interleave_self_int_dim_is_None_shard0(self):
1111- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int((5, 8), 3, [Shard(0)])
1112- 
1113- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value)
1114- output = torch.repeat_interleave(input_tensor, repeats_value)
1115-
1116- self.assertEqual(output_dtensor.full_tensor(), output)
1117- 
1118- @skipIfUnsupportMultiNPU(4)
1119- @with_comms
1120- def test_torch_repeat_interleave_self_int_dim_is_None_shard1(self):
1121- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int((5, 8), 3, [Shard(1)])
1122- 
1123- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value)
1124- output = torch.repeat_interleave(input_tensor, repeats_value)
1125-
1126- self.assertEqual(output_dtensor.full_tensor(), output)
1127- 
1128- @skipIfUnsupportMultiNPU(4)
1129- @with_comms
1130- def test_torch_repeat_interleave_self_int_dim_is_None_shard0_is_evenly_shardable(self):
1131- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int((8, 5), 3, [Shard(0)])
1132- 
1133- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value)
1134- output = torch.repeat_interleave(input_tensor, repeats_value)
1135-
1136- self.assertEqual(output_dtensor.full_tensor(), output)
1137- 
1138- @skipIfUnsupportMultiNPU(4)
1139- @with_comms
1140- def test_torch_repeat_interleave_self_int_shard0_dim1_is_not_evenly_shardable(self):
1141- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int((5, 5), 3, [Shard(0)])
1142- 
1143- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1144- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1145-
1146- self.assertEqual(output_dtensor.full_tensor(), output)
1147- 
1148- @skipIfUnsupportMultiNPU(4)
1149- @with_comms
1150- def test_torch_repeat_interleave_self_int_shard1_dim1_is_not_evenly_shardable(self):
1151- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int((5, 8), 3, [Shard(1)])
1152- 
1153- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1154- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1155-
1156- self.assertEqual(output_dtensor.full_tensor(), output)
1157- 
1158- @skipIfUnsupportMultiNPU(4)
1159- @with_comms
1160- def test_torch_repeat_interleave_backward_self_int_replicate_dim1(self):
1161- sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
1162- 
1163- for size in sizes:
1164- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int(size, 3, [Replicate()])
1165- 
1166- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1167- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1168- 
1169- grad_tensor = torch.randn(output.size(), device="npu")
1170- grad_dtensor = distribute_tensor(grad_tensor, input_dtensor.mesh, [Replicate()])
1171- 
1172- output_dtensor.backward(grad_dtensor)
1173- output.backward(grad_tensor)
1174- self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
1175- 
1176- @skipIfUnsupportMultiNPU(4)
1177- @with_comms
1178- def test_torch_repeat_interleave_backward_self_int_replicate_shard0_dim1(self):
1179- sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
1180- 
1181- for size in sizes:
1182- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int(size, 3, [Replicate()])
1183- 
1184- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1185- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1186- 
1187- grad_tensor = torch.randn(output.size(), device="npu")
1188- grad_dtensor = distribute_tensor(grad_tensor, input_dtensor.mesh, [Shard(0)])
1189- 
1190- output_dtensor.backward(grad_dtensor)
1191- output.backward(grad_tensor)
1192- self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
1193- 
1194- @skipIfUnsupportMultiNPU(4)
1195- @with_comms
1196- def test_torch_repeat_interleave_backward_self_int_shard1_replicate_dim1(self):
1197- sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
1198- 
1199- for size in sizes:
1200- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int(size, 3, [Shard(1)])
1201- 
1202- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1203- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1204- 
1205- grad_tensor = torch.randn(output.size(), device="npu")
1206- grad_dtensor = distribute_tensor(grad_tensor, input_dtensor.mesh, [Replicate()])
1207- 
1208- output_dtensor.backward(grad_dtensor)
1209- output.backward(grad_tensor)
1210- self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
1211- 
1212- @skipIfUnsupportMultiNPU(4)
1213- @with_comms
1214- def test_torch_repeat_interleave_backward_self_int_replicate_dim_None(self):
1215- sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
1216- 
1217- for size in sizes:
1218- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int(size, 3, [Replicate()])
1219- 
1220- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value)
1221- output = torch.repeat_interleave(input_tensor, repeats_value)
1222- 
1223- grad_tensor = torch.randn(output.size(), device="npu")
1224- grad_dtensor = distribute_tensor(grad_tensor, input_dtensor.mesh, [Replicate()])
1225- 
1226- output_dtensor.backward(grad_dtensor)
1227- output.backward(grad_tensor)
1228- self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
1229- 
1230- @skipIfUnsupportMultiNPU(4)
1231- @with_comms
1232- def test_torch_repeat_interleave_backward_self_int_shard00_dim1(self):
1233- sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
1234- 
1235- for size in sizes:
1236- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int(size, 3, [Shard(0)])
1237- 
1238- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1239- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1240- 
1241- grad_tensor = torch.randn(output.size(), device="npu")
1242- grad_dtensor = distribute_tensor(grad_tensor, input_dtensor.mesh, [Shard(0)])
1243- 
1244- output_dtensor.backward(grad_dtensor)
1245- output.backward(grad_tensor)
1246- self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
1247- 
1248- @skipIfUnsupportMultiNPU(4)
1249- @with_comms
1250- def test_torch_repeat_interleave_backward_self_int_shard01_dim1(self):
1251- sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
1252- 
1253- for size in sizes:
1254- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int(size, 3, [Shard(0)])
1255- 
1256- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1257- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1258- 
1259- grad_tensor = torch.randn(output.size(), device="npu")
1260- grad_dtensor = distribute_tensor(grad_tensor, input_dtensor.mesh, [Shard(1)])
1261- 
1262- output_dtensor.backward(grad_dtensor)
1263- output.backward(grad_tensor)
1264- self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
1265- 
1266- @skipIfUnsupportMultiNPU(4)
1267- @with_comms
1268- def test_torch_repeat_interleave_backward_self_int_shard10_dim1(self):
1269- sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
1270- 
1271- for size in sizes:
1272- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int(size, 3, [Shard(1)])
1273- 
1274- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1275- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1276- 
1277- grad_tensor = torch.randn(output.size(), device="npu")
1278- grad_dtensor = distribute_tensor(grad_tensor, input_dtensor.mesh, [Shard(0)])
1279- 
1280- output_dtensor.backward(grad_dtensor)
1281- output.backward(grad_tensor)
1282- self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
1283- 
1284- @skipIfUnsupportMultiNPU(4)
1285- @with_comms
1286- def test_torch_repeat_interleave_backward_self_int_shard11_dim1(self):
1287- sizes = [(2, 2), (5, 5), (5, 8), (8, 5), (8, 8)]
1288- 
1289- for size in sizes:
1290- input_tensor, repeats_value, input_dtensor = self.generate_data_repeat_interleave_self_int(size, 3, [Shard(1)])
1291- 
1292- output_dtensor = torch.repeat_interleave(input_dtensor, repeats_value, dim=1)
1293- output = torch.repeat_interleave(input_tensor, repeats_value, dim=1)
1294- 
1295- grad_tensor = torch.randn(output.size(), device="npu")
1296- grad_dtensor = distribute_tensor(grad_tensor, input_dtensor.mesh, [Shard(1)])
1297- 
1298- output_dtensor.backward(grad_dtensor)
1299- output.backward(grad_tensor)
1300- self.assertEqual(input_dtensor.grad.full_tensor(), input_tensor.grad)
1301- 
1302 250 
1303if __name__ == "__main__":251if __name__ == "__main__":
1304 run_tests()252 run_tests()