Conversation
| config = self.original_model.config | ||
| start = getattr(config, "decoder_start_token_id", None) | ||
| if start is None: | ||
| start = getattr(config, "bos_token_id", None) |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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): |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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( |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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() |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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().
|
Thanks for the great work as always @emerardd |
Description
Fix encoder-decoder streaming generation so it uses the same seq2seq state as
generate(), rather than silently entering the decoder-only loop.On
devat11224618,generate_stream()recognizes encoder-decoder models during tokenization but passesis_encoder_decoder=False,encoder_input=None, anddecoder_tokens=Noneto_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 againstgenerate(): streamed output has seven tokens instead of four, and the cached batched T5 case raises a tensor broadcast error.Changes
generate(), preserving the existing start-token fallback chain and configured forced BOS token.generate(). Do not introduce encoder-decoder KV-cache support in this fix.generate()limitation, with guidance to usehf_generate().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
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
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.