Jump to content

Connect SuperML | Leeroopedia MCP: Equip your AI agents with best practices, code verification, and debugging knowledge. Powered by Leeroo — building Organizational Superintelligence. Contact us at founders@leeroo.com.

Implementation:Predibase Lorax Seq2Seq LM

From Leeroopedia
Revision as of 16:21, 16 February 2026 by Admin (talk | contribs) (Auto-imported from implementations/Predibase_Lorax_Seq2Seq_LM.md)
(diff) ← Older revision | Latest revision (diff) | Newer revision → (diff)


Knowledge Sources
Domains Model_Architecture, Inference
Last Updated 2026-02-08 00:00 GMT

Overview

Implements the sequence-to-sequence language model inference wrapper for encoder-decoder architectures (e.g., T5, BART), including padded batch management with separate encoder and decoder states and autoregressive token generation within the LoRax server.

Description

This module provides the standard (non-flash-attention) encoder-decoder model wrapper using HuggingFace AutoModelForSeq2SeqLM. It handles the dual-state nature of seq2seq models: encoder hidden states and decoder past key-values.

Key classes:

  • Seq2SeqLMBatch (extends Batch, dataclass) - Manages batch state for encoder-decoder inference:
    • Encoder state: input_ids, attention_mask, encoder_last_hidden_state.
    • Decoder state: decoder_input_ids, decoder_attention_mask, all_decoder_input_ids.
    • Tracks 4-tuple past key values per layer: decoder self-attention keys/values and encoder cross-attention keys/values.
    • from_pb - Constructs the batch from protobuf, initializing the decoder with BOS token and the encoder with tokenized inputs.
    • filter - Removes completed requests, slicing both encoder and decoder tensors. Handles the 4-element past key values (2 decoder + 2 encoder per layer).
    • concatenate - Merges multiple batches by padding both encoder and decoder past key value caches to matching dimensions, maintaining separate shapes for decoder self-attention and encoder cross-attention.
  • Seq2SeqLM (extends Model) - The model wrapper that:
    • Loads models via AutoModelForSeq2SeqLM.from_pretrained with support for bitsandbytes 8-bit quantization and multi-GPU device mapping.
    • Sets bos_token_id from decoder_start_token_id in the model config.
    • forward - Runs the full encoder-decoder forward pass, returning logits, encoder last hidden state, and updated past key values.
    • generate_token - Autoregressive generation step that:
      • Wraps encoder_last_hidden_state in a list (Transformers internal requirement).
      • Applies token choosers and stopping criteria.
      • Manages decoder attention mask updates with right padding offset.
      • Sets input_ids to None after prefill (encoder only runs once).
      • Returns generations with sharding across world size.

Usage

Seq2SeqLM is the inference wrapper for encoder-decoder models like T5 and BART within the LoRax server. It is instantiated by the model registry when the loaded model is identified as a seq2seq architecture. Unlike CausalLM, it maintains both encoder and decoder states and handles the encoder-decoder cross-attention cache.

Code Reference

Source Location

  • Repository: Predibase_Lorax
  • File: server/lorax_server/models/seq2seq_lm.py
  • Lines: 1-716

Signature

@dataclass
class Seq2SeqLMBatch(Batch):
    batch_id: int
    requests: List[generate_pb2.Request]
    requests_idx_mapping: Dict[int, int]
    input_ids: Optional[torch.Tensor]
    attention_mask: torch.Tensor
    decoder_input_ids: torch.Tensor
    decoder_attention_mask: Optional[torch.Tensor]
    encoder_last_hidden_state: Optional[torch.Tensor]
    all_decoder_input_ids: List[torch.Tensor]
    past_key_values: Optional[List[Tuple]]
    input_lengths: List[int]
    decoder_input_lengths: List[int]
    prefix_offsets: List[int]
    read_offsets: List[int]
    next_token_choosers: List[NextTokenChooser]
    stopping_criterias: List[StoppingCriteria]
    max_input_length: int
    max_decoder_input_length: int
    padding_right_offset: int
    max_tokens: int

    @classmethod
    def from_pb(cls, pb, tokenizer, tokenizers, processor, config, dtype, device) -> "Seq2SeqLMBatch":
        ...
    def filter(self, request_ids: List[int]) -> Optional["Seq2SeqLMBatch"]:
        ...
    @classmethod
    def concatenate(cls, batches: List["Seq2SeqLMBatch"]) -> "Seq2SeqLMBatch":
        ...

class Seq2SeqLM(Model):
    def __init__(
        self,
        model_id: str,
        revision: Optional[str] = None,
        quantize: Optional[str] = None,
        compile: bool = False,
        dtype: Optional[torch.dtype] = None,
        trust_remote_code: bool = False,
    ):
        ...
    @property
    def batch_type(self) -> Type[Seq2SeqLMBatch]:
        ...
    def forward(self, input_ids, attention_mask, decoder_input_ids, decoder_attention_mask,
                encoder_last_hidden_state, past_key_values=None) -> Tuple[torch.Tensor, torch.Tensor, List[Tuple]]:
        ...
    def generate_token(self, batch: Seq2SeqLMBatch) -> Tuple[List[Generation], Optional[Seq2SeqLMBatch]]:
        ...

Import

from lorax_server.models.seq2seq_lm import Seq2SeqLM, Seq2SeqLMBatch

I/O Contract

Inputs

Name Type Required Description
model_id str Yes HuggingFace model identifier (e.g., "t5-base", "facebook/bart-large")
revision Optional[str] No Model revision/commit hash
quantize Optional[str] No Quantization method (e.g., "bitsandbytes")
compile bool No Whether to compile the model (not supported, logged as skip)
dtype Optional[torch.dtype] No Model dtype (defaults to float16 on GPU, float32 on CPU)
trust_remote_code bool No Whether to trust remote code in model loading

Outputs

Name Type Description
generations List[Generation] Generated token information for each request in the batch
next_batch Optional[Seq2SeqLMBatch] Updated batch for next generation step (None if all requests complete)

Usage Examples

# Internal LoRax server usage
from lorax_server.models.seq2seq_lm import Seq2SeqLM

# Instantiated by model registry for encoder-decoder models
# model = Seq2SeqLM(model_id="t5-base", dtype=torch.float16)

Related Pages

Page Connections

Double-click a node to navigate. Hold to expand connections.
Principle
Implementation
Heuristic
Environment