CoolFace
Apppublic

bohraanuj23/SeverityAnalysis

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
model.py116 linesDownload Raw Back to root
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