Commit 179df148 authored by Kirill Milintsevich's avatar Kirill Milintsevich
Browse files

Old code clean-up

parent 46f90861
Loading
Loading
Loading
Loading
+1 −14
Original line number Diff line number Diff line
@@ -12,19 +12,6 @@ from hdsc.modules import ClassificationHead, SentenceLevelEncoder
PAD_ID = 0


class ModelOutput(NamedTuple):
    pred_binary: torch.Tensor
    pred_binary_final: Optional[torch.Tensor] = None
    pred_regression: Optional[torch.Tensor] = None
    pred_multilabel: Optional[torch.Tensor] = None
    attn_binary: Optional[torch.Tensor] = None
    attn_regression: Optional[torch.Tensor] = None
    word_attns: Optional[torch.Tensor] = None
    word_conicity: Optional[torch.Tensor] = None
    sent_conicity: Optional[torch.Tensor] = None
    chunk_hidden_states: Optional[torch.Tensor] = None


class PHQTotalMulticlassAttentionModelBERT(nn.Module):
    """Hierarchical Attention Classification model with BERT in the word level.

@@ -112,7 +99,7 @@ class PHQTotalMulticlassAttentionModelBERT(nn.Module):
        self,
        inputs: Dict[str, torch.Tensor],
        return_attn: bool = False,
    ) -> ModelOutput:
    ) -> torch.Tensor:
        model_output = self.encoder(inputs["input_ids"], inputs["attention_mask"])
        sentence_embeddings = self.mean_pooling(model_output, inputs["attention_mask"])
        word_outputs = torch.split(sentence_embeddings, inputs["text_lens"].tolist())
+1 −1
Original line number Diff line number Diff line
import json
import itertools
import json
from pathlib import Path
from typing import Dict, List, Optional