import torch import torch.nn as nn class ConvBlock(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.block = nn.Sequential( nn.Conv2d( in_channels, out_channels, kernel_size=3, padding=1 ), nn.ReLU(inplace=True), nn.Conv2d( out_channels, out_channels, kernel_size=3, padding=1 ), nn.ReLU(inplace=True) ) def forward(self, x): return self.block(x) class Encoder(nn.Module): def __init__(self): super().__init__() self.stage_1 = ConvBlock( in_channels=3, out_channels=9 ) self.stage_2 = ConvBlock( in_channels=9, out_channels=27 ) self.stage_3 = nn.Conv2d( in_channels=27, out_channels=3, kernel_size=3, padding=1 ) def forward(self, x): x = self.stage_1(x) x = self.stage_2(x) x = self.stage_3(x) return x class Decoder(nn.Module): def __init__(self): super().__init__() self.upsample_block = nn.Upsample( size=(640, 1024), mode="bilinear", align_corners=False ) def forward(self, x): return self.upsample_block(x) class ImageRegressionNet(nn.Module): def __init__(self): super().__init__() self.encoder = Encoder() self.decoder = Decoder() self.encoder_decoder = nn.Sequential( self.encoder, self.decoder ) def forward(self, x): return self.encoder_decoder(x)