|
19 | 19 |
|
20 | 20 | import nemo.collections.asr.modules.ggemm_transformer_encoder as ggemm_module |
21 | 21 | import nemo.collections.asr.modules.transformer_encoder as transformer_module |
22 | | -from nemo.collections.asr.modules.ggemm_transformer_encoder import GGEMMTransformerEncoder, _ragged_grouped_mm |
23 | 22 | from nemo.collections.asr.parts.packed_sequence import PackedEncoderActivations, pack_encoder_output |
24 | 23 | from tests.collections.asr.test_parallel_expert_encoder import ( |
25 | 24 | _MEL_FEATURES, |
@@ -462,7 +461,7 @@ def forward_sequence_packed(self, audio_signal, length, bypass_pre_encode=False) |
462 | 461 | return PackedEncoderActivations(audio_signal, length, cu_seqlens, int(length.max())) |
463 | 462 |
|
464 | 463 | expert = PackedOnly() |
465 | | - container = GGEMMTransformerEncoder({'custom': expert}) |
| 464 | + container = ggemm_module.GGEMMTransformerEncoder({'custom': expert}) |
466 | 465 | data = torch.randn(3, 4) |
467 | 466 | output = container.forward_all_sequence_packed(data, torch.tensor([3])) |
468 | 467 | assert output['custom'].data is data |
@@ -496,7 +495,7 @@ def test_ragged_grouped_mm_backward_handles_empty_expert_and_sum_loss(offset_val |
496 | 495 | x = torch.randn(7, 16, device='cuda', dtype=torch.bfloat16, requires_grad=True) |
497 | 496 | weights = [torch.randn(16, 8, device='cuda', requires_grad=True) for _ in range(3)] |
498 | 497 | offsets = torch.tensor(offset_values, device='cuda', dtype=torch.int32) |
499 | | - output = _ragged_grouped_mm(x, offsets, weights) |
| 498 | + output = ggemm_module._ragged_grouped_mm(x, offsets, weights) |
500 | 499 | output.sum().backward() |
501 | 500 |
|
502 | 501 | assert x.grad is not None and torch.isfinite(x.grad).all() |
|
0 commit comments