Spaces:
Runtime error
Runtime error
| from vnet_blocks import * | |
| class VNetLight(BaseModel): | |
| """ | |
| A lighter version of Vnet that skips down_tr256 and up_tr256 in oreder to reduce time and space complexity | |
| """ | |
| def __init__(self, elu=True, in_channels=1, classes=4): | |
| super(VNetLight, self).__init__() | |
| self.classes = classes | |
| self.in_channels = in_channels | |
| self.in_tr = InputTransition(in_channels, elu) | |
| self.down_tr32 = DownTransition(16, 1, elu) | |
| self.down_tr64 = DownTransition(32, 2, elu) | |
| self.down_tr128 = DownTransition(64, 3, elu, dropout=True) | |
| self.up_tr128 = UpTransition(128, 128, 2, elu, dropout=True) | |
| self.up_tr64 = UpTransition(128, 64, 1, elu) | |
| self.up_tr32 = UpTransition(64, 32, 1, elu) | |
| self.out_tr = OutputTransition(32, classes, elu) | |
| def forward(self, x): | |
| out16 = self.in_tr(x) | |
| out32 = self.down_tr32(out16) | |
| out64 = self.down_tr64(out32) | |
| out128 = self.down_tr128(out64) | |
| out = self.up_tr128(out128, out64) | |
| out = self.up_tr64(out, out32) | |
| out = self.up_tr32(out, out16) | |
| out = self.out_tr(out) | |
| return out | |
| def test(self,device='cpu'): | |
| input_tensor = torch.rand(1, self.in_channels, 32, 32, 32) | |
| ideal_out = torch.rand(1, self.classes, 32, 32, 32) | |
| out = self.forward(input_tensor) | |
| assert ideal_out.shape == out.shape | |
| summary(self.to(torch.device(device)), (self.in_channels, 32, 32, 32),device=device) | |
| print("Vnet Light test is complete") |