import torch import torch.nn as nn from transformers import PreTrainedModel from .configuration_alpha3d import Alpha3DConfig class Alpha3DModel(PreTrainedModel): config_class = Alpha3DConfig def __init__(self, config): super().__init__(config) self.config = config layers = [] input_dim = 5 prev_dim = input_dim for h_dim in config.hidden_dims: layers.append(nn.Linear(prev_dim, h_dim)) layers.append(nn.BatchNorm1d(h_dim)) layers.append(nn.ReLU()) prev_dim = h_dim layers.append(nn.Linear(prev_dim, config.num_points * 6)) self.net = nn.Sequential(*layers) def forward(self, x): out = self.net(x) return out.view(-1, self.config.num_points, 6)