nemo_automodel.components.models.common.packing

View as Markdown

Flash Attention packing support via monkey-patching.

When attn_implementation="flash_attention_2" and neat packing is enabled, the collater produces an indexed attention mask [B, S] where each position contains the 1-based document index (0 = padding). For example::

[1, 1, 2, 2, 2, 0] # 2 tokens in doc 1, 3 in doc 2, 1 padding

To make HuggingFace’s flash attention path use flash_attn_varlen_func with per-document cu_seqlens, we monkey-patch two functions:

  1. transformers.modeling_flash_attention_utils._get_unpad_data — extracts per-document sequence lengths from the indexed mask and builds cu_seqlens.
  2. transformers.models.qwen3_vl.modeling_qwen3_vl.create_causal_mask — returns the 2D indexed mask as-is, bypassing 4D mask creation.

This is the same approach used by LlamaFactory.

Module Contents

Functions

NameDescription
_model_attn_implementationReturn the packing-relevant attention backend an already-built model runs with.
_passthrough_create_causal_maskReplacement for create_causal_mask that passes through packed masks.
_patch_preprocess_mask_arguments_for_packingKeep indexed packing masks intact for the supported FA2 path.
configure_packingApply monkey-patches for packed-sequence training with flash attention.
get_attn_implementationDetermine the attention backend from model config.
get_seqlens_in_batchExtract per-document sequence lengths from an indexed attention mask.
get_unpad_dataPrepare indices and cu_seqlens for flash_attn_varlen_func.
is_indexed_packed_maskReturn True iff attention_mask is an Automodel-style indexed packing mask.

Data

_FLASH_ATTN_IMPLEMENTATIONS

_PACKING_PATCH_MODULES

logger

API

nemo_automodel.components.models.common.packing._model_attn_implementation(
model
) -> str | None

Return the packing-relevant attention backend an already-built model runs with.

model.config._attn_implementation is a Transformers dispatch key, whose vocabulary is wider than the mask layouts packing knows about: when flash attention is requested but only the kernels package provides it, Transformers records a kernels-hub id instead of the mainline name. Those ids are mapped back so a model genuinely running varlen flash attention is packed as such. Any key that still names no known layout yields None, leaving the caller on the configured value.

nemo_automodel.components.models.common.packing._passthrough_create_causal_mask(
config = None,
input_embeds = None,
inputs_embeds = None,
attention_mask = None,
cache_position = None,
past_key_values = None,
position_ids = None,
kwargs = {}
)

Replacement for create_causal_mask that passes through packed masks.

Flash attention (FA2/FA3/FA4) handles masking internally, so always pass through. For other backends, pass through packed masks but delegate normal 2D masks to HF.

nemo_automodel.components.models.common.packing._patch_preprocess_mask_arguments_for_packing() -> None

Keep indexed packing masks intact for the supported FA2 path.

Transformers 5.x preprocesses 2D attention masks before dispatching attention. For flash attention this can coerce integer indexed masks (1, 2, ... per packed document) to bool masks, losing the document boundaries that get_unpad_data needs. Preserve indexed 2D masks for FA2 so the patched flash-attention path can derive per-document cu_seqlens. Validate the private Transformers contract before installing the shim so an incompatible dependency fails instead of silently enabling cross-document attention.

nemo_automodel.components.models.common.packing.configure_packing(
attn_implementation: str = 'sdpa'
) -> None

Apply monkey-patches for packed-sequence training with flash attention.

Only patches when attn_implementation is a flash-attention variant (flash_attention_2 / flash_attention_3 / flash_attention_4); transformers routes all three through the same varlen wrapper, so the _get_unpad_data patch applies uniformly.

Parameters:

attn_implementation
strDefaults to 'sdpa'

The attention implementation used by the model.

nemo_automodel.components.models.common.packing.get_attn_implementation(
cfg_model,
model = None
) -> str

Determine the attention backend from model config.

Custom models store it in backend.attn; HF models use attn_implementation.

Parameters:

cfg_model

Model config node, which records what was requested.

model
Defaults to None

Optional already-built model, preferred over cfg_model for HF models because it records what was actually resolved. Model construction may pick a different backend than the config asks for: packed runs are force-switched onto flash attention (_apply_preload_overrides), an unavailable backend is downgraded on retry, and an omitted key defaults to flash attention rather than to sdpa. None of those are written back to the config. An HF model configured with te reports sdpa here, which is what it runs with TE attention injected on top.

nemo_automodel.components.models.common.packing.get_seqlens_in_batch(
attention_mask: torch.Tensor
) -> torch.Tensor

Extract per-document sequence lengths from an indexed attention mask.

Example::

>>> get_seqlens_in_batch(torch.tensor([[1, 1, 2, 2, 2, 0], … [1, 2, 2, 3, 3, 3]])) tensor([2, 3, 1, 2, 3])

Parameters:

attention_mask
torch.Tensor

[B, S] integer tensor where each position contains the 1-based document index (0 = padding).

Returns: torch.Tensor

1D tensor of all individual document lengths across the batch.

nemo_automodel.components.models.common.packing.get_unpad_data(
attention_mask: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor, int]

Prepare indices and cu_seqlens for flash_attn_varlen_func.

This is a drop-in replacement for transformers.modeling_flash_attention_utils._get_unpad_data that handles indexed attention masks (values 1, 2, 3, …) instead of binary (0/1) masks. Each unique non-zero value is treated as a separate document, so flash_attn_varlen_func applies causal attention within each document without cross-document attention.

Example::

>>> get_unpad_data(torch.tensor([[1, 1, 2, 2, 2, 0], … [1, 2, 2, 3, 3, 3]])) (tensor([0, 1, 2, 3, 4, 6, 7, 8, 9, 10, 11]), tensor([ 0, 2, 5, 6, 8, 11], dtype=torch.int32), 3)

Returns: torch.Tensor

Indices of non-padding tokens from the flattened sequence.

nemo_automodel.components.models.common.packing.is_indexed_packed_mask(
attention_mask: torch.Tensor | None
) -> bool

Return True iff attention_mask is an Automodel-style indexed packing mask.

The Automodel neat_packed_vlm_collater (and the LLM equivalent) encode packed-sample boundaries by marking document i (1-based) with the integer i and using 0 for padding (e.g. [1, 1, 1, 2, 2, 3, 3, 0, 0]). Any value greater than 1 is therefore a sufficient signal that two or more documents are packed into the same row. A standard 0/1 attention mask never has values > 1.

nemo_automodel.components.models.common.packing._FLASH_ATTN_IMPLEMENTATIONS = ('flash_attention_2', 'flash_attention_3', 'flash_attention_4')
nemo_automodel.components.models.common.packing._PACKING_PATCH_MODULES = ['transformers.models.llama.modeling_llama', 'transformers.models.qwen3.modeling...
nemo_automodel.components.models.common.packing.logger = logging.getLogger(__name__)