Skip to content

fix(bridge): preserve seq2seq state in streaming generation - #1825

Merged
jlarson4 merged 2 commits into
TransformerLensOrg:devfrom
emerardd:fix/encoder-decoder-generate-stream
Sep 30, 2026
Merged

jlarson4 merged 2 commits into
TransformerLensOrg:devfrom
emerardd:fix/encoder-decoder-generate-stream

Conversation

@emerardd

Copy link
Copy Markdown
Contributor

Description

Fix encoder-decoder streaming generation so it uses the same seq2seq state as generate(), rather than silently entering the decoder-only loop.

On dev at 11224618, generate_stream() recognizes encoder-decoder models during tokenization but passes is_encoder_decoder=False, encoder_input=None, and decoder_tokens=None to _generate_tokens(). It grows the encoder prompt and derives new decoder inputs from it on every step. With offline tiny BART/T5 models, all eight combinations of single/batched input and requested caching enabled/disabled fail a greedy comparison against generate(): streamed output has seven tokens instead of four, and the cached batched T5 case raises a tensor broadcast error.

Changes

  • Keep encoder inputs fixed and grow only the decoder sequence.
  • Share decoder-prefix construction with generate(), preserving the existing start-token fallback chain and configured forced BOS token.
  • Include the decoder prefix, rather than the encoder prompt, in the first streamed chunk.
  • Preserve tokenizer-produced encoder attention masks for string/list inputs and honor seq2seq minimum-length and no-repeat-ngram settings.
  • Use the existing uncached seq2seq loop, matching generate(). Do not introduce encoder-decoder KV-cache support in this fix.
  • Explicitly reject seq2seq custom stopping criteria, matching the existing generate() limitation, with guidance to use hf_generate().
  • Document the seq2seq chunk contract in the public method docstring.

Decoder-only streaming behavior and public method signatures are unchanged. This does not attempt to make the native loop support every HuggingFace generation option.

Type of change

  • Bug fix (non-breaking change which fixes an issue)

Validation

Offline tiny BART/T5 regressions cover exact greedy stream/generate parity, single/batched input, both requested cache settings, fixed encoder/growing decoder forward traces, string output with unequal padded inputs, encoder masks, EOS termination, forced BOS/minimum length/no-repeat-ngram defaults, decoder-start fallback cases, and unsupported stopping criteria.

The affected existing streaming, seq2seq-loss, generation-capability, and stopping-criteria tests were also run: 100 passed, 6 skipped, 2 warnings. All 30 new regressions passed without skips. The six skips are existing downloaded-model tests whose models are unavailable in the local offline cache. Full-suite testing was not rerun for this PR.

Environment: Windows, Python 3.12.10, Torch 2.11.0+cpu, Transformers 5.13.0; installed from source with uv. Formatting checks (pycln/isort/Black) passed for both changed files, and mypy . passed for 385 source files.

Checklist

  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes
  • I have not rewritten tests relating to key interfaces which would affect backward compatibility

The unchecked unit-test item reflects that only the affected surface was run, not the complete unit suite. The test run reports dependency SWIG deprecation warnings; these are not hidden by this PR.

@jlarson4 jlarson4 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks putting this bug fix together @emerardd. Routing the stream through generate()'s seq2seq state is the right fix. A couple small comments below

config = self.original_model.config
start = getattr(config, "decoder_start_token_id", None)
if start is None:
start = getattr(config, "bos_token_id", None)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Moving this out of generate() dropped the comments explaining the bos → eos fallback (IndicBART) and why a forced BOS follows decoder_start (target-language selection). Could you bring those comments into the helper? It is nice to keep some indicator of WHY to provide future contirbutors context.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added the rationale to the shared setup helper: the decoder-start fallback follows BOS then EOS (including IndicBART), and the forced BOS comes after decoder_start for target-language selection.

# stop_strings / stopping_criteria: build the combined criteria list (validates
# tokenizer for stop_strings). generate_stream only runs the decoder-only text
# path, so no path guards are needed here.
if _encdec_early and (stop_strings is not None or stopping_criteria is not None):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This raises on no-op inputs like StoppingCriteriaList() or stop_strings=[], which generate() accepts on the same model. Can the guard key off the resolved criteria, as generate() does?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The guard now checks the resolved criteria. Empty stop_strings and an empty StoppingCriteriaList are accepted; I added parity tests for both, while active criteria remain explicitly unsupported for seq2seq streaming.

decoder_tokens = None
stream_prefix = input_tokens
if _encdec_early:
forced_bos = getattr(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

generate() and generate_stream() still each read forced_bos_token_id and min_length off generation_config, and drift between duplicated setup like this is what caused the bug. Could the shared helper own those reads too?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The shared _encoder_decoder_setup() now reads both forced_bos_token_id and min_length from generation_config, with generate() still able to pass its explicit forced-BOS override. Both generation paths consume the same setup result.

chunks = list(seq2seq_bridge.generate_stream(source, max_tokens_per_yield=2, **options))
torch.testing.assert_close(torch.cat(chunks, dim=1), expected, rtol=0, atol=0)
assert chunks[0][0, 1] == 8
assert (expected[:, 2:4] != 2).all()

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This holds even if the stream ignores min_length, since the tiny models never pick token 2 at these steps. Could you set up a case where EOS would otherwise fire early, the way test_stream_stops_at_eos picks its EOS?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added a discriminating min_length regression: it first probes the model's greedy first token, uses that token as EOS, then checks that min_length suppresses it at the first step and that the stream exactly matches generate().

@jlarson4

Copy link
Copy Markdown
Collaborator

Thanks for the great work as always @emerardd

@jlarson4
jlarson4 merged commit 02a7f5a into TransformerLensOrg:dev Sep 30, 2026
27 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants