classification¶
Sequence classification implementation.
Classes
Collation function for sequence classification. |
|
|
Sequence classification head. |
|
A transformer encoder with a sequence classification head. |
- class undertale.models.classification.ClassificationCollator¶
Bases:
objectCollation function for sequence classification.
Stacks
tokensandmasktensors and gathers integerlabelvalues into a 1-D tensor.
- class undertale.models.classification.ClassificationHead(hidden_dimensions: int, classes: int)¶
Bases:
ModuleSequence classification head.
A single linear projection from hidden state space to class logits.
- Parameters:
hidden_dimensions – The size of the hidden state space.
classes – The number of output classes.
- forward(state: Tensor) Tensor¶
Project hidden state to class logits.
- Parameters:
state – Pooled encoder hidden state.
- Returns:
A tensor of class logits.
- class undertale.models.classification.InstructionTraceTransformerEncoderForSequenceClassification(depth: int, hidden_dimensions: int, vocab_size: int, sequence_length: int, heads: int, intermediate_dimensions: int, next_token_id: int, classes: int, dropout: float, eps: float, lr: float = 0.0001, warmup: float = 0.025, class_weights: List[float] | None = None)¶
Bases:
LightningModule,ModuleA transformer encoder with a sequence classification head.
- Parameters:
depth – The number of stacked transformer layers.
hidden_dimensions – The size of the hidden state space.
vocab_size – The size of the vocabulary.
sequence_length – The fixed size of the input vector.
heads – The number of attention heads.
intermediate_dimensions – The size of the intermediate state space.
next_token_id – The ID of the special
NEXTtoken.classes – The number of output classes.
dropout – Dropout probability.
eps – Layer normalization stabalization parameter.
lr – Peak learning rate reached after warmup.
warmup – Fraction of total steps used for linear warmup.
- forward(state: Tensor, mask: Tensor | None = None) Tensor¶
Encode and classify the input sequence.
- Parameters:
state – The tokenized input state tensor.
mask – Optional attention mask.
- Returns:
A tensor of class logits derived from masked, mean-pooled hidden state.