From caa7dfb972b0c8cc1c5c878defcde790139eaf32 Mon Sep 17 00:00:00 2001 From: Shiyi Zheng Date: Thu, 3 Sep 2026 17:23:56 +0800 Subject: [PATCH 1/3] Add bounded translation evaluation support --- src/winml/modelkit/eval/__init__.py | 3 + src/winml/modelkit/eval/evaluate.py | 2 + .../modelkit/eval/metrics/translation.py | 53 ++++ .../modelkit/eval/translation_evaluator.py | 213 +++++++++++++++ .../modelkit/models/winml/encoder_decoder.py | 23 +- src/winml/modelkit/utils/eval_utils.py | 70 +++++ tests/unit/eval/test_translation_evaluator.py | 245 ++++++++++++++++++ 7 files changed, 608 insertions(+), 1 deletion(-) create mode 100644 src/winml/modelkit/eval/metrics/translation.py create mode 100644 src/winml/modelkit/eval/translation_evaluator.py create mode 100644 tests/unit/eval/test_translation_evaluator.py diff --git a/src/winml/modelkit/eval/__init__.py b/src/winml/modelkit/eval/__init__.py index 435601a15..ae00585df 100644 --- a/src/winml/modelkit/eval/__init__.py +++ b/src/winml/modelkit/eval/__init__.py @@ -41,6 +41,7 @@ from .text_classification_evaluator import WinMLTextClassificationEvaluator from .text_generation_evaluator import WinMLTextGenerationEvaluator from .token_classification_evaluator import WinMLTokenClassificationEvaluator + from .translation_evaluator import WinMLTranslationEvaluator from .zero_shot_classification_evaluator import WinMLZeroShotClassificationEvaluator from .zero_shot_image_classification_evaluator import WinMLZeroShotImageClassificationEvaluator @@ -73,6 +74,7 @@ "WinMLTokenClassificationEvaluator": ( ".token_classification_evaluator:WinMLTokenClassificationEvaluator" ), + "WinMLTranslationEvaluator": ".translation_evaluator:WinMLTranslationEvaluator", "WinMLZeroShotClassificationEvaluator": ( ".zero_shot_classification_evaluator:WinMLZeroShotClassificationEvaluator" ), @@ -140,6 +142,7 @@ def __dir__() -> list[str]: "WinMLTextClassificationEvaluator", "WinMLTextGenerationEvaluator", "WinMLTokenClassificationEvaluator", + "WinMLTranslationEvaluator", "WinMLZeroShotClassificationEvaluator", "WinMLZeroShotImageClassificationEvaluator", "evaluate", diff --git a/src/winml/modelkit/eval/evaluate.py b/src/winml/modelkit/eval/evaluate.py index 42a9ab8a5..0de3d7a3e 100644 --- a/src/winml/modelkit/eval/evaluate.py +++ b/src/winml/modelkit/eval/evaluate.py @@ -84,6 +84,8 @@ def _select_model_loader(config: WinMLEvaluationConfig) -> _ModelLoaderKind: "winml.modelkit.eval.image_feature_extraction_evaluator:WinMLImageFeatureExtractionEvaluator", "image-to-text": "winml.modelkit.eval.image_to_text_evaluator:WinMLImageToTextEvaluator", + "translation": + "winml.modelkit.eval.translation_evaluator:WinMLTranslationEvaluator", "fill-mask": "winml.modelkit.eval.fill_mask_evaluator:WinMLFillMaskEvaluator", "zero-shot-classification": diff --git a/src/winml/modelkit/eval/metrics/translation.py b/src/winml/modelkit/eval/metrics/translation.py new file mode 100644 index 000000000..9963ddee5 --- /dev/null +++ b/src/winml/modelkit/eval/metrics/translation.py @@ -0,0 +1,53 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + +"""Corpus-level machine-translation metrics.""" + +from __future__ import annotations + +from typing import Any + + +class TranslationMetric: + """Aggregate corpus SacreBLEU-13a and chrF2 on a 0-100 scale.""" + + def __init__(self) -> None: + self._predictions: list[str] = [] + self._references: list[list[str]] = [] + + def update(self, prediction: str, references: str | list[str]) -> None: + """Record one prediction and one or more non-empty references.""" + refs = [references] if isinstance(references, str) else references + cleaned = [reference.strip() for reference in refs if reference and reference.strip()] + if not cleaned: + raise ValueError("at least one non-empty translation reference is required") + self._predictions.append((prediction or "").strip()) + self._references.append(cleaned) + + def compute(self) -> dict[str, Any]: + """Return corpus scores with names that state variant and scale.""" + if not self._predictions: + return { + "sacrebleu_13a_0_100": None, + "chrf2_0_100": None, + "n_samples": 0, + } + + from torchmetrics.text import CHRFScore, SacreBLEUScore + + sacrebleu = SacreBLEUScore(tokenize="13a")( + self._predictions, + self._references, + ) + chrf = CHRFScore(n_char_order=6, n_word_order=0, beta=2.0)( + self._predictions, + self._references, + ) + + return { + "sacrebleu_13a_0_100": round(float(sacrebleu) * 100, 4), + "chrf2_0_100": round(float(chrf) * 100, 4), + "n_samples": len(self._predictions), + } diff --git a/src/winml/modelkit/eval/translation_evaluator.py b/src/winml/modelkit/eval/translation_evaluator.py new file mode 100644 index 000000000..a942145b8 --- /dev/null +++ b/src/winml/modelkit/eval/translation_evaluator.py @@ -0,0 +1,213 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + +"""Evaluator for text-to-text machine translation models.""" + +from __future__ import annotations + +import logging +from collections.abc import Mapping +from typing import TYPE_CHECKING, Any + +from .base_evaluator import WinMLEvaluator + + +if TYPE_CHECKING: + from datasets import Dataset + + from ..models.winml.composite_model import WinMLCompositeModel + from .config import DatasetConfig, WinMLEvaluationConfig + + +logger = logging.getLogger(__name__) + + +def _positive_int(mapping: Mapping[str, str], name: str, default: int) -> int: + from ..utils.eval_utils import DatasetValidationError + + raw_value = mapping.get(name, str(default)) + try: + value = int(raw_value) + except (TypeError, ValueError) as error: + raise DatasetValidationError(f"{name} must be a positive integer") from error + if value < 1: + raise DatasetValidationError(f"{name} must be a positive integer") + return value + + +class WinMLTranslationEvaluator(WinMLEvaluator): + """Evaluate translations with bounded generation and corpus metrics.""" + + def __init__( + self, + config: WinMLEvaluationConfig, + model: WinMLCompositeModel, + ) -> None: + from ..utils.eval_utils import DatasetValidationError, get_default + + mapping = config.dataset.columns_mapping + self._source_col = mapping.get( + "source_column", + get_default("translation", "source_column") or "translation", + ) + self._reference_col = mapping.get( + "reference_column", + get_default("translation", "reference_column") or "translation", + ) + self._source_lang = mapping.get("source_lang") + self._target_lang = mapping.get("target_lang") + self._tokenizer_source_lang = mapping.get("tokenizer_source_lang", self._source_lang) + self._tokenizer_target_lang = mapping.get("tokenizer_target_lang", self._target_lang) + self._source_prefix = mapping.get("source_prefix", "") + max_source_tokens = _positive_int(mapping, "max_source_tokens", 128) + max_new_tokens = _positive_int(mapping, "max_new_tokens", 64) + num_beams = _positive_int(mapping, "num_beams", 1) + num_return_sequences = _positive_int(mapping, "num_return_sequences", 1) + if num_beams != 1 or num_return_sequences != 1: + raise DatasetValidationError( + "translation evaluation requires num_beams=1 and num_return_sequences=1 " + "for the static batch-one generation contract" + ) + + super().__init__(config, model) + self._pipeline_kwargs: dict[str, Any] = { + "do_sample": False, + "max_new_tokens": max_new_tokens, + "num_beams": 1, + "num_return_sequences": 1, + "truncation": True, + } + if self._tokenizer_source_lang: + self._pipeline_kwargs["src_lang"] = self._tokenizer_source_lang + if self._tokenizer_target_lang: + self._pipeline_kwargs["tgt_lang"] = self._tokenizer_target_lang + + max_encoder_length = getattr(model, "max_encoder_length", None) + if isinstance(max_encoder_length, int) and max_encoder_length > 0: + max_source_tokens = min(max_source_tokens, max_encoder_length) + if self.pipe.tokenizer is not None: + self.pipe.tokenizer.model_max_length = max_source_tokens + + max_decode_length = getattr(model, "max_decode_length", None) + if isinstance(max_decode_length, int): + if max_decode_length <= 1: + raise DatasetValidationError( + "decoder static cache must leave capacity for at least one generated token" + ) + self._pipeline_kwargs["max_new_tokens"] = min( + max_new_tokens, + max_decode_length - 1, + ) + self.pipe.generation_config.max_length = max_decode_length + self.pipe.generation_config.max_new_tokens = None + self.pipe.generation_config.num_beams = 1 + self.pipe.generation_config.num_return_sequences = 1 + self.pipe.generation_config.do_sample = False + + def align_labels(self, dataset: Dataset, ds_config: DatasetConfig) -> Dataset: + """Return free-text references unchanged.""" + return dataset + + @staticmethod + def _extract_text(value: Any, language: str | None, field: str) -> str: + """Extract one text value from a flat string or language-keyed mapping.""" + from ..utils.eval_utils import DatasetValidationError + + if isinstance(value, Mapping): + if not language: + language_option = "source_lang" if field == "source" else "target_lang" + raise DatasetValidationError( + f"{field} contains a translation dict; provide --column " + f"{language_option}=" + ) + if language not in value: + raise DatasetValidationError( + f"{field} translation dict has no language key '{language}'; " + f"available keys: {sorted(str(key) for key in value)}" + ) + value = value[language] + if not isinstance(value, str) or not value.strip(): + raise DatasetValidationError(f"{field} must resolve to a non-empty string") + return value.strip() + + def _extract_references(self, value: Any) -> str | list[str]: + """Extract one or more reference translations.""" + if isinstance(value, list): + if not value: + from ..utils.eval_utils import DatasetValidationError + + raise DatasetValidationError("reference must contain at least one translation") + return [self._extract_text(item, self._target_lang, "reference") for item in value] + return self._extract_text(value, self._target_lang, "reference") + + @staticmethod + def _prediction_text(output: Any) -> str: + """Normalize Hugging Face translation pipeline output shapes.""" + from ..utils.eval_utils import DatasetValidationError + + if isinstance(output, list): + if not output: + raise DatasetValidationError("pipeline returned no translations") + output = output[0] + if isinstance(output, Mapping): + prediction = output.get("translation_text", output.get("generated_text", "")) + if isinstance(prediction, str) and prediction.strip(): + return prediction.strip() + raise DatasetValidationError("pipeline returned an empty translation") + if isinstance(output, str) and output.strip(): + return output.strip() + raise DatasetValidationError("pipeline returned an unsupported translation result") + + def _extract_sample(self, sample: Any) -> tuple[str, str | list[str]]: + """Extract and validate source/reference values from one dataset row.""" + from ..utils.eval_utils import DatasetValidationError + + if not isinstance(sample, Mapping): + raise DatasetValidationError( + f"dataset row must be a mapping, got {type(sample).__name__}" + ) + source = self._extract_text(sample.get(self._source_col), self._source_lang, "source") + references = self._extract_references(sample.get(self._reference_col)) + return f"{self._source_prefix}{source}", references + + def compute(self) -> dict[str, Any]: + """Translate each valid row and compute corpus-level metrics.""" + from tqdm.auto import tqdm + + from ..utils.eval_utils import DatasetValidationError + from .metrics.translation import TranslationMetric + + metric = TranslationMetric() + skipped = 0 + first_error: str | None = None + attempted = 0 + + for sample in tqdm(self.data, desc="Evaluating", unit="sample"): + attempted += 1 + try: + source, references = self._extract_sample(sample) + except DatasetValidationError as error: + logger.warning("Translation sample rejected: %s", error) + first_error = first_error or str(error) + skipped += 1 + continue + + output = self.pipe(source, **self._pipeline_kwargs) + try: + prediction = self._prediction_text(output) + metric.update(prediction, references) + except (DatasetValidationError, ValueError) as error: + logger.warning("Translation output rejected: %s", error) + first_error = first_error or str(error) + skipped += 1 + + result = metric.compute() + if result["n_samples"] == 0: + detail = f": {first_error}" if first_error else "" + raise DatasetValidationError(f"No valid translation samples were evaluated{detail}") + result["attempted"] = attempted + result["evaluated"] = result["n_samples"] + result["skipped"] = skipped + return result diff --git a/src/winml/modelkit/models/winml/encoder_decoder.py b/src/winml/modelkit/models/winml/encoder_decoder.py index 6388a42f1..9b538e248 100644 --- a/src/winml/modelkit/models/winml/encoder_decoder.py +++ b/src/winml/modelkit/models/winml/encoder_decoder.py @@ -193,6 +193,12 @@ def __init__( enc_expected = dict( zip(enc_io.get("input_names", []), enc_io.get("input_shapes", []), strict=False) ) + input_ids_shape = enc_expected.get("input_ids", []) + self._max_enc = ( + input_ids_shape[1] + if len(input_ids_shape) > 1 and isinstance(input_ids_shape[1], int) + else None + ) self._encoder_input_names = frozenset(enc_expected) # Wrap encoder with auto-padding so all callsites just use self._encoder(...) self._encoder = self._EncoderWithInputPadding(raw_encoder, enc_expected) @@ -203,7 +209,12 @@ def __init__( ) # Max decode length and KV dtype from decoder ONNX metadata - self._max_dec = self._dec_expected["past_0_key"][2] + max_decode_length = self._dec_expected["past_0_key"][2] + if not isinstance(max_decode_length, int) or max_decode_length <= 0: + raise ValueError( + "Decoder input 'past_0_key' must have a positive static cache length" + ) + self._max_dec = max_decode_length self._num_kv_layers = sum( 1 for n in self._dec_expected if n.startswith("past_") and n.endswith("_key") ) @@ -222,6 +233,16 @@ def __init__( _np_dtype = dec_type_map["past_0_key"] self._kv_dtype = torch.from_numpy(np.zeros(1, dtype=_np_dtype)).dtype + @property + def max_encoder_length(self) -> int | None: + """Return the encoder's static token capacity, or ``None`` when dynamic.""" + return self._max_enc + + @property + def max_decode_length(self) -> int: + """Return the decoder's static output/cache capacity.""" + return self._max_dec + # ----- Encoder ----- class _EncoderWithInputPadding(torch.nn.Module): diff --git a/src/winml/modelkit/utils/eval_utils.py b/src/winml/modelkit/utils/eval_utils.py index f85c672b4..82091fc8c 100644 --- a/src/winml/modelkit/utils/eval_utils.py +++ b/src/winml/modelkit/utils/eval_utils.py @@ -243,6 +243,75 @@ class TaskSchema: roles=("encoder", "decoder"), ) +_TRANSLATION_SCHEMA = TaskSchema( + columns=( + SchemaItem( + "source_column", + "source text, or a translation dict containing the source language key", + default="translation", + remap_hint="", + ), + SchemaItem( + "reference_column", + "reference text(s), or a translation dict containing the target language key", + default="translation", + remap_hint="", + ), + ), + params=( + SchemaItem( + "source_lang", + "source-language key when source_column contains translation dicts", + remap_hint="", + ), + SchemaItem( + "target_lang", + "target-language key when reference_column contains translation dicts", + remap_hint="", + ), + SchemaItem( + "tokenizer_source_lang", + "source-language identifier for multilingual tokenizers; defaults to source_lang", + remap_hint="", + ), + SchemaItem( + "tokenizer_target_lang", + "target-language identifier for multilingual tokenizers; defaults to target_lang", + remap_hint="", + ), + SchemaItem( + "source_prefix", + "optional task prefix prepended before tokenization", + remap_hint="", + ), + SchemaItem( + "max_source_tokens", + "maximum source tokens, clamped to the encoder capacity", + default="128", + remap_hint="", + ), + SchemaItem( + "max_new_tokens", + "maximum generated tokens, clamped to the decoder capacity", + default="64", + remap_hint="", + ), + SchemaItem( + "num_beams", + "beam count; static batch-one models require 1", + default="1", + remap_hint="<1>", + ), + SchemaItem( + "num_return_sequences", + "returned sequences per source; static batch-one models require 1", + default="1", + remap_hint="<1>", + ), + ), + roles=("encoder", "decoder"), +) + _FILL_MASK_SCHEMA = TaskSchema( columns=( SchemaItem( @@ -461,6 +530,7 @@ class TaskSchema: "sentence-similarity": _FEATURE_EXTRACTION_SCHEMA, "image-feature-extraction": _IMAGE_FEATURE_EXTRACTION_SCHEMA, "image-to-text": _IMAGE_TO_TEXT_SCHEMA, + "translation": _TRANSLATION_SCHEMA, "fill-mask": _FILL_MASK_SCHEMA, "zero-shot-classification": _ZERO_SHOT_CLASSIFICATION_SCHEMA, "zero-shot-image-classification": _ZERO_SHOT_IMAGE_CLASSIFICATION_SCHEMA, diff --git a/tests/unit/eval/test_translation_evaluator.py b/tests/unit/eval/test_translation_evaluator.py new file mode 100644 index 000000000..2ebd386d3 --- /dev/null +++ b/tests/unit/eval/test_translation_evaluator.py @@ -0,0 +1,245 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + +"""Unit tests for translation evaluation.""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest +from datasets import Dataset + +from winml.modelkit.eval.translation_evaluator import WinMLTranslationEvaluator +from winml.modelkit.utils.eval_utils import DatasetValidationError + + +def make_evaluator(rows, columns_mapping=None, pipeline=None): + from winml.modelkit.eval import DatasetConfig, WinMLEvaluationConfig + + dataset = Dataset.from_list(rows) + model = MagicMock() + model.config.label2id = None + model.max_encoder_length = 512 + model.max_decode_length = 512 + config = WinMLEvaluationConfig( + model_id="example/translation-model", + task="translation", + dataset=DatasetConfig( + path="example/parallel-corpus", + samples=len(rows), + shuffle=False, + columns_mapping=columns_mapping or {}, + ), + ) + with ( + patch("datasets.load_dataset", return_value=dataset), + patch( + "winml.modelkit.eval.base_evaluator.WinMLEvaluator.prepare_pipeline", + return_value=pipeline or MagicMock(), + ), + ): + return WinMLTranslationEvaluator(config, model) + + +class TestTranslationMetric: + def test_perfect_corpus_scores_100(self): + from winml.modelkit.eval.metrics.translation import TranslationMetric + + metric = TranslationMetric() + sentence = "This is a complete test sentence." + metric.update(sentence, sentence) + + assert metric.compute() == { + "sacrebleu_13a_0_100": 100.0, + "chrf2_0_100": 100.0, + "n_samples": 1, + } + + def test_empty_corpus_has_no_scores(self): + from winml.modelkit.eval.metrics.translation import TranslationMetric + + assert TranslationMetric().compute() == { + "sacrebleu_13a_0_100": None, + "chrf2_0_100": None, + "n_samples": 0, + } + + def test_multiple_references_use_corpus_metrics(self): + from winml.modelkit.eval.metrics.translation import TranslationMetric + + metric = TranslationMetric() + prediction = "A sufficiently long reference sentence is here." + metric.update(prediction, ["Another complete reference is here.", prediction]) + + assert metric.compute()["sacrebleu_13a_0_100"] == 100.0 + + def test_empty_references_are_rejected(self): + from winml.modelkit.eval.metrics.translation import TranslationMetric + + with pytest.raises(ValueError, match="at least one non-empty"): + TranslationMetric().update("prediction", []) + + +class TestTranslationEvaluator: + def test_nested_translation_uses_explicit_direction_and_bounded_generation(self): + evaluator = make_evaluator( + [{"translation": {"source": "Bonjour le monde", "target": "Hello world"}}], + {"source_lang": "source", "target_lang": "target"}, + ) + assert evaluator.pipe.tokenizer.model_max_length == 128 + assert evaluator.pipe.generation_config.max_new_tokens is None + evaluator.pipe = MagicMock(return_value=[{"translation_text": "Hello world"}]) + + result = evaluator.compute() + + evaluator.pipe.assert_called_once_with( + "Bonjour le monde", + do_sample=False, + max_new_tokens=64, + num_beams=1, + num_return_sequences=1, + truncation=True, + src_lang="source", + tgt_lang="target", + ) + assert result["chrf2_0_100"] == 100.0 + assert result["attempted"] == result["evaluated"] == 1 + assert result["skipped"] == 0 + + def test_explicit_caps_clamp_to_static_model_capacities(self): + pipeline = MagicMock() + evaluator = make_evaluator( + [{"source": "texte", "reference": "text"}], + { + "source_column": "source", + "reference_column": "reference", + "max_source_tokens": "1024", + "max_new_tokens": "1024", + }, + pipeline, + ) + evaluator.pipe = MagicMock(return_value={"generated_text": "text"}) + + evaluator.compute() + + assert pipeline.tokenizer.model_max_length == 512 + assert evaluator.pipe.call_args.kwargs["max_new_tokens"] == 511 + + @pytest.mark.parametrize( + ("name", "value"), + [("max_source_tokens", "0"), ("max_new_tokens", "bad")], + ) + def test_invalid_token_bounds_are_rejected(self, name, value): + with pytest.raises(DatasetValidationError, match=f"{name} must be a positive integer"): + make_evaluator( + [{"source": "texte", "reference": "text"}], + {"source_column": "source", "reference_column": "reference", name: value}, + ) + + @pytest.mark.parametrize("name", ["num_beams", "num_return_sequences"]) + def test_static_batch_contract_rejects_generation_fanout(self, name): + with pytest.raises(DatasetValidationError, match="static batch-one"): + make_evaluator( + [{"source": "texte", "reference": "text"}], + {"source_column": "source", "reference_column": "reference", name: "2"}, + ) + + def test_pipeline_without_tokenizer_is_supported(self): + pipeline = MagicMock() + pipeline.tokenizer = None + + evaluator = make_evaluator( + [{"source": "texte", "reference": "text"}], + {"source_column": "source", "reference_column": "reference"}, + pipeline, + ) + + assert evaluator.pipe.tokenizer is None + + def test_flat_columns_and_multiple_references(self): + evaluator = make_evaluator( + [{"source": "Une phrase source", "references": ["A source sentence", "A sentence"]}], + {"source_column": "source", "reference_column": "references"}, + ) + evaluator.pipe = MagicMock(return_value={"generated_text": "A source sentence"}) + + assert evaluator.compute()["n_samples"] == 1 + + def test_missing_language_key_is_skipped_with_exact_accounting(self): + evaluator = make_evaluator( + [ + {"translation": {"source": "Valide", "target": "Valid"}}, + {"translation": {"source": "Invalide", "other": "Invalid"}}, + ], + {"source_lang": "source", "target_lang": "target"}, + ) + evaluator.pipe = MagicMock(return_value=[{"translation_text": "Valid"}]) + + result = evaluator.compute() + + assert result["attempted"] == 2 + assert result["evaluated"] == 1 + assert result["skipped"] == 1 + evaluator.pipe.assert_called_once() + + def test_nested_translation_requires_explicit_direction(self): + evaluator = make_evaluator( + [{"translation": {"source": "Bonjour", "target": "Hello"}}] + ) + + with pytest.raises(DatasetValidationError, match="provide --column source_lang"): + evaluator.compute() + + def test_tokenizer_languages_and_prefix_are_independent_of_dataset_keys(self): + evaluator = make_evaluator( + [{"translation": {"dataset_source": "Bonjour", "dataset_target": "Hello"}}], + { + "source_lang": "dataset_source", + "target_lang": "dataset_target", + "tokenizer_source_lang": "tokenizer_source", + "tokenizer_target_lang": "tokenizer_target", + "source_prefix": "translate: ", + }, + ) + evaluator.pipe = MagicMock(return_value=[{"translation_text": "Hello"}]) + + evaluator.compute() + + assert evaluator.pipe.call_args.args == ("translate: Bonjour",) + assert evaluator.pipe.call_args.kwargs["src_lang"] == "tokenizer_source" + assert evaluator.pipe.call_args.kwargs["tgt_lang"] == "tokenizer_target" + + def test_runtime_failure_propagates(self): + evaluator = make_evaluator( + [{"source": "un", "reference": "one"}], + {"source_column": "source", "reference_column": "reference"}, + ) + evaluator.pipe = MagicMock(side_effect=RuntimeError("decoder failed")) + + with pytest.raises(RuntimeError, match="decoder failed"): + evaluator.compute() + + @pytest.mark.parametrize("output", [[], [{}], [{"translation_text": ""}]]) + def test_invalid_output_fails_closed_when_no_rows_are_evaluated(self, output): + evaluator = make_evaluator( + [{"source": "texte", "reference": "text"}], + {"source_column": "source", "reference_column": "reference"}, + ) + evaluator.pipe = MagicMock(return_value=output) + + with pytest.raises(DatasetValidationError, match="No valid translation samples"): + evaluator.compute() + + +class TestTranslationRegistration: + def test_registry_and_schema_are_present(self): + from winml.modelkit.eval import WinMLEvaluationConfig, get_evaluator_class + from winml.modelkit.utils.eval_utils import TASK_SCHEMAS + + assert get_evaluator_class(WinMLEvaluationConfig(task="translation")) is ( + WinMLTranslationEvaluator + ) + assert TASK_SCHEMAS["translation"].roles == ("encoder", "decoder") From 92387df1fbe50b628ac38574cdcb0fa992247088 Mon Sep 17 00:00:00 2001 From: Shiyi Zheng Date: Thu, 3 Sep 2026 17:24:41 +0800 Subject: [PATCH 2/3] Refresh opus-mt-en-ru eager attention recipes --- .../cpu/cpu/translation_fp16_decoder_config.json | 5 ++++- .../cpu/cpu/translation_fp16_encoder_config.json | 5 ++++- .../cpu/cpu/translation_fp32_decoder_config.json | 5 ++++- .../cpu/cpu/translation_fp32_encoder_config.json | 5 ++++- 4 files changed, 16 insertions(+), 4 deletions(-) diff --git a/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp16_decoder_config.json b/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp16_decoder_config.json index 1fd73fdc0..2cba18f03 100644 --- a/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp16_decoder_config.json +++ b/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp16_decoder_config.json @@ -36,7 +36,10 @@ {"name": "present_3_key"}, {"name": "present_3_value"}, {"name": "present_4_key"}, {"name": "present_4_value"}, {"name": "present_5_key"}, {"name": "present_5_value"} - ] + ], + "compatibility": { + "transformers_attention": "eager" + } }, "optim": { "clamp_constant_values": true, diff --git a/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp16_encoder_config.json b/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp16_encoder_config.json index 45840beea..c3752b775 100644 --- a/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp16_encoder_config.json +++ b/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp16_encoder_config.json @@ -13,7 +13,10 @@ {"name": "input_ids", "dtype": "int32", "shape": [1, 512], "value_range": [0, 62518]}, {"name": "attention_mask", "dtype": "int32", "shape": [1, 512], "value_range": [0, 2]} ], - "output_tensors": [{"name": "encoder_hidden_states"}] + "output_tensors": [{"name": "encoder_hidden_states"}], + "compatibility": { + "transformers_attention": "eager" + } }, "optim": { "clamp_constant_values": true, diff --git a/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp32_decoder_config.json b/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp32_decoder_config.json index 86090caf3..a0c773b0e 100644 --- a/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp32_decoder_config.json +++ b/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp32_decoder_config.json @@ -36,7 +36,10 @@ {"name": "present_3_key"}, {"name": "present_3_value"}, {"name": "present_4_key"}, {"name": "present_4_value"}, {"name": "present_5_key"}, {"name": "present_5_value"} - ] + ], + "compatibility": { + "transformers_attention": "eager" + } }, "optim": { "clamp_constant_values": true, diff --git a/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp32_encoder_config.json b/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp32_encoder_config.json index 977e346e8..53536778d 100644 --- a/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp32_encoder_config.json +++ b/examples/recipes/Helsinki-NLP_opus-mt-en-ru/cpu/cpu/translation_fp32_encoder_config.json @@ -13,7 +13,10 @@ {"name": "input_ids", "dtype": "int32", "shape": [1, 512], "value_range": [0, 62518]}, {"name": "attention_mask", "dtype": "int32", "shape": [1, 512], "value_range": [0, 2]} ], - "output_tensors": [{"name": "encoder_hidden_states"}] + "output_tensors": [{"name": "encoder_hidden_states"}], + "compatibility": { + "transformers_attention": "eager" + } }, "optim": { "clamp_constant_values": true, From 3f5f321e4721e3022f7f50514e8cfcfb9ad4570d Mon Sep 17 00:00:00 2001 From: Shiyi Zheng Date: Thu, 3 Sep 2026 18:31:25 +0800 Subject: [PATCH 3/3] Preserve encoder mask alignment during generation --- .../modelkit/models/winml/encoder_decoder.py | 8 ++ .../winml/test_composite_from_pretrained.py | 101 ++++++++++++++++++ 2 files changed, 109 insertions(+) diff --git a/src/winml/modelkit/models/winml/encoder_decoder.py b/src/winml/modelkit/models/winml/encoder_decoder.py index 9b538e248..fae99eb70 100644 --- a/src/winml/modelkit/models/winml/encoder_decoder.py +++ b/src/winml/modelkit/models/winml/encoder_decoder.py @@ -407,6 +407,14 @@ def _run_decoder( dtype=torch.int64, ), ) + attention_mask = runtime_feeds.get("attention_mask") + expected_attention_mask = self._dec_expected.get("attention_mask") + if isinstance(attention_mask, torch.Tensor) and expected_attention_mask is not None: + runtime_feeds["attention_mask"] = pad_inputs( + {"attention_mask": attention_mask}, + {"attention_mask": expected_attention_mask}, + mode="right", + )["attention_mask"] for i in range(self._num_kv_layers): layer = cache._layer(i) runtime_feeds[f"past_{i}_key"] = cast("torch.Tensor", layer.keys).detach() diff --git a/tests/unit/models/winml/test_composite_from_pretrained.py b/tests/unit/models/winml/test_composite_from_pretrained.py index 1f12839ba..a2ca4db32 100644 --- a/tests/unit/models/winml/test_composite_from_pretrained.py +++ b/tests/unit/models/winml/test_composite_from_pretrained.py @@ -355,6 +355,107 @@ def _advance(outputs): assert cache.step == 3 +def test_encoder_decoder_preserves_encoder_alignment_across_independent_rows() -> None: + from winml.modelkit.models.winml.encoder_decoder import WinMLEncoderDecoderModel + + class _SourceConditionedDecoder: + def __init__(self) -> None: + self.io_config = { + "input_names": [ + "decoder_input_ids", + "encoder_hidden_states", + "attention_mask", + "decoder_attention_mask", + "cache_position", + "past_0_key", + "past_0_value", + ], + "input_shapes": [ + [1, 1], + [1, 4, 2], + [1, 4], + [1, 4], + [1], + [1, 2, 4, 2], + [1, 2, 4, 2], + ], + "input_types": [ + np.int64, + np.float32, + np.int64, + np.int64, + np.int64, + np.float32, + np.float32, + ], + } + self.attention_masks: list[torch.Tensor] = [] + + def __call__(self, **feeds): + attention_mask = feeds["attention_mask"] + self.attention_masks.append(attention_mask.clone()) + source_token = int( + (feeds["encoder_hidden_states"][:, :, 0] * attention_mask).sum().item() + ) + logits = torch.zeros(1, 1, 4) + logits[0, 0, source_token] = 1 + present = torch.full((1, 2, 1, 2), float(source_token)) + return { + "logits": logits, + "present_0_key": present, + "present_0_value": present.clone(), + } + + def make_cache(): + cache = MagicMock() + cache.step = 0 + cache.layers = [ + SimpleNamespace( + keys=torch.zeros(1, 2, 4, 2), + values=torch.zeros(1, 2, 4, 2), + ) + ] + cache.build_decoder_mask.return_value = torch.tensor([[1, 0, 0, 0]]) + cache.get_query_cache_position.return_value = torch.tensor([0]) + + def _advance(outputs): + cache.step += outputs["present_0_key"].shape[2] + + cache.update_all_layers.side_effect = _advance + return cache + + decoder = _SourceConditionedDecoder() + model = WinMLEncoderDecoderModel( + {"encoder": SimpleNamespace(io_config={}), "decoder": decoder}, + MagicMock(is_encoder_decoder=True), + ) + first_cache = make_cache() + second_cache = make_cache() + rows = [ + (torch.tensor([[[1.0, 0.0], [0.0, 0.0], [0.0, 0.0], [0.0, 0.0]]]), 1), + (torch.tensor([[[2.0, 0.0], [0.0, 0.0], [0.0, 0.0], [0.0, 0.0]]]), 2), + ] + + generated_tokens = [] + with patch.object(model, "_resolve_cache", side_effect=[first_cache, second_cache]): + for hidden_state, _ in rows: + result = model.forward( + encoder_outputs=BaseModelOutput(last_hidden_state=hidden_state), + decoder_input_ids=torch.tensor([[0]]), + attention_mask=torch.tensor([[1]]), + ) + generated_tokens.append(int(result.logits[:, -1, :].argmax(dim=-1)[0])) + + assert [mask.tolist() for mask in decoder.attention_masks] == [ + [[1, 0, 0, 0]], + [[1, 0, 0, 0]], + ] + assert generated_tokens == [1, 2] + assert first_cache.step == second_cache.step == 1 + assert first_cache.update_all_layers.call_count == 1 + assert second_cache.update_all_layers.call_count == 1 + + def test_static_cache_rejects_out_of_range_query_positions() -> None: from winml.modelkit.models.winml.kv_cache import WinMLStaticCache