MCPcopy Create free account
hub / github.com/apple/axlearn / _pad

Method _pad

axlearn/common/decoder.py:461–473  ·  view source on GitHub ↗

Accept token IDs input tensor and pad if necessary to max_sequence_length.

(prefix: Tensor, *, max_sequence_length: int, pad_id: int)

Source from the content-addressed store, hash-verified

459 dtype=prefix.dtype,
460 ),
461 ],
462 axis=1,
463 )
464
465
466# TODO(gyin): Add unittest for Decoder forward
467class Decoder(BaseLayer):
468 """Construct a decoder transformer to output hidden states and logits based on lm head."""
469
470 @config_class
471 class Config(BaseLayer.Config):
472 """Configures Decoder."""
473
474 # DEPRECATED, because `attention_mask` uses quadratic memory, even with Flash Attention.
475 # Please use `attention.mask`, which constructs the mask procedurally.
476 attention_mask: Optional[AttentionLogitBiasLayer.Config] = None

Callers 2

beam_search_decodeMethod · 0.95
sample_decodeMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected