bohraanuj23/SeverityAnalysis
0
1import torch2import torch.nn as nn3import pytorch_lightning as pl4 5 6class ResidualBlock(nn.Module):7 def __init__(self, in_features, out_features, dropout=0.2):8 super().__init__()9 self.fc1 = nn.Linear(in_features, out_features)10 self.relu = nn.ReLU()11 self.dropout = nn.Dropout(dropout)12 self.fc2 = nn.Linear(out_features, out_features)13 14 self.projection = (15 nn.Linear(in_features, out_features)16 if in_features != out_features17 else nn.Identity()18 )19 20 def forward(self, x):21 residual = self.projection(x)22 out = self.fc1(x)23 out = self.relu(out)24 out = self.dropout(out)25 out = self.fc2(out)26 return out + residual27 28 29class DualEncoderModel(pl.LightningModule):30 def __init__(31 self,32 lab_cont_dim,33 lab_cat_dims,34 conv_cont_dim,35 conv_cat_dims,36 embedding_dim,37 num_classes,38 lr=1e-3,39 ):40 super().__init__()41 self.save_hyperparameters()42 43 # Lab continuous44 self.lab_cont_encoder = (45 nn.Sequential(ResidualBlock(lab_cont_dim, 64), ResidualBlock(64, 64))46 if lab_cont_dim > 047 else None48 )49 50 # Lab categorical51 self.lab_cat_embeddings = nn.ModuleList(52 [nn.Embedding(dim + 1, embedding_dim) for dim in lab_cat_dims]53 )54 55 # Conversation continuous56 self.conv_cont_encoder = (57 nn.Sequential(ResidualBlock(conv_cont_dim, 64), ResidualBlock(64, 64))58 if conv_cont_dim > 059 else None60 )61 62 # Conversation categorical63 self.conv_cat_embeddings = nn.ModuleList(64 [nn.Embedding(dim + 1, embedding_dim) for dim in conv_cat_dims]65 )66 67 # Calculate total input dimension to classifier68 total_dim = 069 if self.lab_cont_encoder:70 total_dim += 6471 if lab_cat_dims:72 total_dim += embedding_dim * len(lab_cat_dims)73 if self.conv_cont_encoder:74 total_dim += 6475 if conv_cat_dims:76 total_dim += embedding_dim * len(conv_cat_dims)77 78 self.classifier = nn.Sequential(79 nn.Linear(total_dim, 128),80 nn.ReLU(),81 nn.Dropout(0.3),82 nn.Linear(128, num_classes),83 )84 85 def forward(self, lab_cont, lab_cat, conv_cont, conv_cat):86 embeddings = []87 88 # Lab continuous89 if self.lab_cont_encoder and lab_cont.nelement() > 0:90 embeddings.append(self.lab_cont_encoder(lab_cont))91 92 # Lab categorical93 if self.lab_cat_embeddings and lab_cat.nelement() > 0:94 embeddings.extend(95 [96 emb(torch.clamp(lab_cat[:, i], min=0))97 for i, emb in enumerate(self.lab_cat_embeddings)98 ]99 )100 101 # Conv continuous102 if self.conv_cont_encoder and conv_cont.nelement() > 0:103 embeddings.append(self.conv_cont_encoder(conv_cont))104 105 # Conv categorical106 if self.conv_cat_embeddings and conv_cat.nelement() > 0:107 embeddings.extend(108 [109 emb(torch.clamp(conv_cat[:, i], min=0))110 for i, emb in enumerate(self.conv_cat_embeddings)111 ]112 )113 114 fused = torch.cat(embeddings, dim=1)115 return self.classifier(fused)116 