(self)
| 553 | self.assertEqual(m.buffer_name, buffer3) |
| 554 | |
| 555 | def test_get_buffer(self): |
| 556 | m = nn.Module() |
| 557 | buffer1 = torch.randn(2, 3) |
| 558 | buffer2 = torch.randn(4, 5) |
| 559 | m.register_buffer('foo', buffer1) |
| 560 | m.register_buffer('bar', buffer2) |
| 561 | self.assertEqual(buffer1, m.get_buffer('foo')) |
| 562 | self.assertEqual(buffer2, m.get_buffer('bar')) |
| 563 | |
| 564 | def test_get_buffer_from_submodules(self): |
| 565 | class MyModule(nn.Module): |
nothing calls this directly
no test coverage detected