Implementation:Predibase Lorax Seq2Seq LM
| 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.
- Encoder state:
- Seq2SeqLM (extends
Model) - The model wrapper that:- Loads models via
AutoModelForSeq2SeqLM.from_pretrainedwith support for bitsandbytes 8-bit quantization and multi-GPU device mapping. - Sets
bos_token_idfromdecoder_start_token_idin 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_statein a list (Transformers internal requirement). - Applies token choosers and stopping criteria.
- Manages decoder attention mask updates with right padding offset.
- Sets
input_idsto None after prefill (encoder only runs once). - Returns generations with sharding across world size.
- Wraps
- Loads models via
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)