from __future__ import annotations

import torch
from torch import nn


class BiGRUSeqClassifier(nn.Module):
    """Embedding -> BiGRU -> Linear (per-step logits)."""

    def __init__(self, vocab_size: int, hidden_size: int, num_layers: int = 1, emb_dim: int | None = None, dropout: float = 0.1):
        super().__init__()
        emb_dim = emb_dim or hidden_size
        self.embedding = nn.Embedding(vocab_size, emb_dim, padding_idx=0)
        self.rnn = nn.GRU(
            input_size=emb_dim,
            hidden_size=hidden_size,
            num_layers=num_layers,
            batch_first=True,
            bidirectional=True,
            dropout=dropout if num_layers > 1 else 0.0,
        )
        self.head = nn.Linear(hidden_size * 2, vocab_size)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        emb = self.embedding(x)
        out, _ = self.rnn(emb)
        return self.head(out)

