admesh/agentic-intent-classifier
245
1from __future__ import annotations2 3from dataclasses import dataclass4 5import torch6from torch import nn7from transformers import AutoModel8 9 10@dataclass(frozen=True)11class MultiTaskLabelSizes:12 intent_type: int13 intent_subtype: int14 decision_phase: int15 16 17class MultiTaskIntentModel(nn.Module):18 def __init__(self, base_model_name: str, label_sizes: MultiTaskLabelSizes):19 super().__init__()20 self.base_model_name = base_model_name21 self.encoder = AutoModel.from_pretrained(base_model_name)22 hidden_size = int(self.encoder.config.hidden_size)23 self.dropout = nn.Dropout(float(getattr(self.encoder.config, "seq_classif_dropout", 0.2)))24 self.intent_type_head = nn.Linear(hidden_size, label_sizes.intent_type)25 self.intent_subtype_head = nn.Linear(hidden_size, label_sizes.intent_subtype)26 self.decision_phase_head = nn.Linear(hidden_size, label_sizes.decision_phase)27 28 def forward(self, input_ids: torch.Tensor, attention_mask: torch.Tensor) -> dict[str, torch.Tensor]:29 outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask)30 pooled = outputs.last_hidden_state[:, 0]31 pooled = self.dropout(pooled)32 return {33 "intent_type_logits": self.intent_type_head(pooled),34 "intent_subtype_logits": self.intent_subtype_head(pooled),35 "decision_phase_logits": self.decision_phase_head(pooled),36 }37 