Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .gitattributes
Original file line number Diff line number Diff line change
@@ -1 +1,2 @@
*.ipynb eol=lf
*.sh eol=lf
2 changes: 1 addition & 1 deletion local_check.sh
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ if [ "$#" -ne 0 ]; then
exit 2
fi

poetry install
poetry install --all-extras

if [ "$agent_strict" = true ]; then
echo "=================== comment hygiene ================="
Expand Down
12 changes: 9 additions & 3 deletions machine/jobs/huggingface/hugging_face_nmt_model_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,8 @@
import datasets.utils.logging as datasets_logging
import transformers.utils.logging as transformers_logging
from transformers import AutoConfig, AutoModelForSeq2SeqLM, HfArgumentParser, PreTrainedModel, Seq2SeqTrainingArguments
from transformers.integrations import ClearMLCallback
from transformers.tokenization_utils import TruncationStrategy
from transformers.integrations.integration_utils import ClearMLCallback
from transformers.tokenization_utils_base import TruncationStrategy

from ...corpora.parallel_text_corpus import ParallelTextCorpus
from ...corpora.text_corpus import TextCorpus
Expand All @@ -26,7 +26,13 @@ def __init__(self, config: Any) -> None:
self._config = config
args = config.huggingface.train_params.to_dict()
args["output_dir"] = str(self._model_dir)
args["overwrite_output_dir"] = True
# Allow group_by_length backwards compatibility. The settings default for train_sampling_strategy is
# group_by_length, so any other value was set explicitly and takes precedence over the legacy option.
group_by_length = args.pop("group_by_length", None)
if group_by_length is not None:
logger.warning("'group_by_length' is deprecated. Use 'train_sampling_strategy' instead.")
if args.get("train_sampling_strategy", "group_by_length") == "group_by_length":
args["train_sampling_strategy"] = "group_by_length" if group_by_length else "random"
# Use "max_steps" from root for backward compatibility
if "max_steps" in self._config.huggingface:
args["max_steps"] = self._config.huggingface.max_steps
Expand Down
1 change: 1 addition & 0 deletions machine/jobs/nmt_build_options.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ class TrainParams(BaseModel):
gradient_accumulation_steps: int | None = None
label_smoothing_factor: float | None = None
group_by_length: bool | None = None
train_sampling_strategy: str | None = None
Comment thread
pmachapman marked this conversation as resolved.
gradient_checkpointing: bool | None = None
lr_scheduler_type: str | None = None
learning_rate: float | None = None
Expand Down
3 changes: 2 additions & 1 deletion machine/jobs/nmt_engine_build_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,7 +149,8 @@ def _translate(
if check_canceled is not None:
check_canceled()
source_segments = [pt_info["translation"] for pt_info in pt_batch]
for pt_info, result in zip(pt_batch, engine.translate_batch(source_segments), strict=True):
t_batch = engine.translate_batch(source_segments)
for pt_info, result in zip(pt_batch, t_batch, strict=True):
pt_info["translation"] = result.translation
pt_info["sequenceConfidence"] = result.sequence_confidence
current_inference_step += len(pt_batch)
Expand Down
2 changes: 1 addition & 1 deletion machine/jobs/settings.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ default:
per_device_train_batch_size: 64
gradient_accumulation_steps: 1
label_smoothing_factor: 0.2
group_by_length: true
train_sampling_strategy: group_by_length
gradient_checkpointing: true
lr_scheduler_type: cosine
learning_rate: 0.0002
Expand Down
10 changes: 8 additions & 2 deletions machine/translation/huggingface/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,14 @@
if not is_torch_available():
raise RuntimeError("torch is not installed.")

from .hugging_face_nmt_engine import HuggingFaceNmtEngine
from .hugging_face_nmt_engine import HuggingFaceNmtEngine, SilTranslationPipeline
from .hugging_face_nmt_model import HuggingFaceNmtModel
from .hugging_face_nmt_model_trainer import HuggingFaceNmtModelTrainer, add_lang_code_to_tokenizer

__all__ = ["add_lang_code_to_tokenizer", "HuggingFaceNmtEngine", "HuggingFaceNmtModel", "HuggingFaceNmtModelTrainer"]
__all__ = [
"add_lang_code_to_tokenizer",
"HuggingFaceNmtEngine",
"HuggingFaceNmtModel",
"HuggingFaceNmtModelTrainer",
"SilTranslationPipeline",
]
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
{
"additional_special_tokens": null,
"extra_special_tokens": null,
"bos_token": "<s>",
"cls_token": "<s>",
"eos_token": "</s>",
Expand Down
158 changes: 111 additions & 47 deletions machine/translation/huggingface/hugging_face_nmt_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,8 @@
import logging
import re
from math import exp, prod
from typing import Collection, Iterable, List, Optional, Sequence, Tuple, Union, cast
from pathlib import Path
from typing import Any, Collection, Iterable, List, Optional, Sequence, Tuple, Union, cast

import torch # pyright: ignore[reportMissingImports]
from sacremoses import MosesPunctNormalizer
Expand All @@ -14,47 +15,61 @@
AutoTokenizer,
M2M100Tokenizer,
NllbTokenizer,
NllbTokenizerFast,
PreTrainedModel,
PreTrainedTokenizer,
PreTrainedTokenizerBase,
PreTrainedTokenizerFast,
TranslationPipeline,
)
from transformers.generation import BeamSearchEncoderDecoderOutput, GreedySearchEncoderDecoderOutput
from transformers.tokenization_utils import BatchEncoding, TruncationStrategy
from transformers.generation.utils import GenerateBeamEncoderDecoderOutput, GenerateEncoderDecoderOutput
from transformers.tokenization_utils_base import BatchEncoding, TruncationStrategy

from ...annotations.range import Range
from ...corpora.aligned_word_pair import AlignedWordPair
from ...utils.typeshed import StrPath
from ..translation_engine import TranslationEngine
from ..translation_result import TranslationResult
from ..translation_result_builder import TranslationResultBuilder
from ..translation_sources import TranslationSources
from ..word_alignment_matrix import WordAlignmentMatrix
from .transformers_compatibility import TranslationPipeline

logger = logging.getLogger(__name__)


class HuggingFaceNmtEngine(TranslationEngine):
def __init__(
self,
model: Union[PreTrainedModel, StrPath, str],
model: Union[PreTrainedModel, Path, str],
oom_batch_size_backoff_mult: float = 1.0,
**pipeline_kwargs,
) -> None:
self._model = model
"""
A model passed as a `PreTrainedModel` stays owned by the caller. When `output_attentions` is enabled, the
engine switches that model to eager attention and restores its original implementation in `close()`.
"""
self._pipeline_kwargs = pipeline_kwargs
if isinstance(self._model, PreTrainedModel):
if self._pipeline_kwargs.get("output_attentions") is None:
self._pipeline_kwargs["output_attentions"] = True
if isinstance(model, PreTrainedModel):
self._model = model
self._model_attn_implementation = model.config._attn_implementation
self._model.eval()
self._is_model_owned = False
else:
model_config = AutoConfig.from_pretrained(str(self._model), label2id={}, id2label={}, num_labels=0)
model_config = AutoConfig.from_pretrained(str(model), label2id={}, id2label={}, num_labels=0)

# Only eager attention returns attention weights. Otherwise, let transformers choose, since forcing "sdpa"
# fails to load architectures without SDPA support (e.g. T5).
attn_implementation = "eager" if self._pipeline_kwargs["output_attentions"] else None

self._model = cast(
PreTrainedModel, AutoModelForSeq2SeqLM.from_pretrained(str(self._model), config=model_config)
PreTrainedModel,
AutoModelForSeq2SeqLM.from_pretrained(
str(model), config=model_config, attn_implementation=attn_implementation
),
)
self._is_model_owned = True
self._tokenizer = AutoTokenizer.from_pretrained(self._model.name_or_path, use_fast=True)
if isinstance(self._tokenizer, (NllbTokenizer, NllbTokenizerFast)):
self._tokenizer = AutoTokenizer.from_pretrained(self._model.name_or_path)
if isinstance(self._tokenizer, NllbTokenizer):
self._mpn = MosesPunctNormalizer()
self._mpn.substitutions = [ # type: ignore
(re.compile(r), sub)
Expand All @@ -70,11 +85,12 @@ def __init__(
src_lang is not None
and tgt_lang is not None
and "prefix" not in self._pipeline_kwargs
and self._model.name_or_path is not None
and (self._model.name_or_path.startswith("t5-") or self._model.name_or_path.startswith("google/mt5-"))
):
self._pipeline_kwargs["prefix"] = f"translate {src_lang} to {tgt_lang}: "
else:
additional_special_tokens = cast(list[str], self._tokenizer.additional_special_tokens or [])
extra_special_tokens = cast(list[str], self._tokenizer.extra_special_tokens or [])
if isinstance(self._tokenizer, M2M100Tokenizer):
src_lang_token = self._tokenizer.lang_code_to_token.get(src_lang) if src_lang is not None else None
tgt_lang_token = self._tokenizer.lang_code_to_token.get(tgt_lang) if tgt_lang is not None else None
Expand All @@ -84,29 +100,33 @@ def __init__(
if (
src_lang is not None
and src_lang_token not in self._tokenizer.added_tokens_encoder
and src_lang_token not in additional_special_tokens
and src_lang_token not in extra_special_tokens
):
raise ValueError(f"The specified model does not support the language code '{src_lang}'")

Comment thread
claude[bot] marked this conversation as resolved.
if (
tgt_lang is not None
and tgt_lang_token not in self._tokenizer.added_tokens_encoder
and tgt_lang_token not in additional_special_tokens
and tgt_lang_token not in extra_special_tokens
):
raise ValueError(f"The specified model does not support the language code '{tgt_lang}'")

self._batch_size = int(self._pipeline_kwargs.pop("batch_size", 1))

self._oom_batch_size_backoff_mult = oom_batch_size_backoff_mult

self._pipeline = _TranslationPipeline(
self._pipeline = SilTranslationPipeline(
model=self._model,
tokenizer=self._tokenizer,
tokenizer=self.tokenizer,
mpn=self._mpn,
batch_size=self._batch_size,
**self._pipeline_kwargs,
)

# Last, so that a constructor that raises leaves the caller's model unchanged.
if not self._is_model_owned and self._pipeline_kwargs["output_attentions"]:
self._model.set_attn_implementation("eager")

@property
def tokenizer(self) -> PreTrainedTokenizer | PreTrainedTokenizerFast:
return self._tokenizer
Expand Down Expand Up @@ -139,7 +159,7 @@ def translate_n_batch(
raise
self._batch_size = max(int(round(self._batch_size * self._oom_batch_size_backoff_mult)), 1)
logger.warning(f"Out of memory error caught. Reducing batch size to {self._batch_size} and retrying.")
self._pipeline = _TranslationPipeline(
self._pipeline = SilTranslationPipeline(
Comment thread
claude[bot] marked this conversation as resolved.
model=self._model,
tokenizer=self._tokenizer,
batch_size=self._batch_size,
Expand Down Expand Up @@ -184,16 +204,19 @@ def close(self) -> None:
del self._pipeline
if self._is_model_owned:
del self._model
elif self._model_attn_implementation is not None:

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

FYI: F12. close() restores the saved implementation even when this engine did not change it, so two engines on one caller model break each other. Repro: engines A and B on an SDPA stas/tiny-m2m_100, both with output_attentions=True. B saves eager, A.close() sets sdpa, then B.translate(...) raises IndexError (via F11). Restoring only when this engine made the switch would avoid it.

# Restore the attn implementation to the model
self._model.set_attn_implementation(self._model_attn_implementation)
gc.collect()
with torch.no_grad():
torch.cuda.empty_cache()


class _TranslationPipeline(TranslationPipeline):
class SilTranslationPipeline(TranslationPipeline):
def __init__(
self,
model: Union[PreTrainedModel, StrPath, str],
tokenizer: Union[PreTrainedTokenizer, PreTrainedTokenizerFast],
model: PreTrainedModel,
tokenizer: PreTrainedTokenizerBase,
batch_size: int,
mpn: Optional[MosesPunctNormalizer] = None,
**kwargs,
Expand Down Expand Up @@ -236,51 +259,43 @@ def preprocess(self, *args, truncation=TruncationStrategy.DO_NOT_TRUNCATE, src_l
return inputs

def _forward(self, model_inputs, **generate_kwargs):
if self.tokenizer is None:
raise RuntimeError("No tokenizer is specified.")
in_b, input_length = model_inputs["input_ids"].shape

input_tokens = model_inputs["input_tokens"]
del model_inputs["input_tokens"]
if hasattr(self.model, "generation_config") and self.model.generation_config is not None:
config = self.model.generation_config
else:
config = self.model.config
generate_kwargs["min_length"] = generate_kwargs.get("min_length", config.min_length)
generate_kwargs["max_length"] = generate_kwargs.get("max_length", config.max_length)
generate_kwargs["output_attentions"] = generate_kwargs.get("output_attentions", True)
self.check_inputs(input_length, generate_kwargs["min_length"], generate_kwargs["max_length"])
output = self.model.generate(
input_tokens = model_inputs.pop("input_tokens")

self.check_inputs(input_length, self.generation_config.min_length, self.generation_config.max_length)
output = cast(Any, self.model).generate(
**model_inputs,
**generate_kwargs,
generation_config=self.generation_config,
output_scores=True,
return_dict_in_generate=True,
)

if isinstance(output, BeamSearchEncoderDecoderOutput):
if isinstance(output, GenerateBeamEncoderDecoderOutput):
output_ids = output.sequences
beam_indices = output.beam_indices
scores = output.scores
assert scores is not None and beam_indices is not None
sequences_scores = output.sequences_scores
attentions = output.cross_attentions
elif isinstance(output, GreedySearchEncoderDecoderOutput):
# Beam search scores are already log probabilities.
normalize_logits = False
Comment thread
ddaspit marked this conversation as resolved.
elif isinstance(output, GenerateEncoderDecoderOutput):
output_ids = output.sequences
beam_indices = None
assert output.scores is not None
scores = output.scores
sequences_scores = None
Comment thread
ddaspit marked this conversation as resolved.
attentions = output.cross_attentions
# Greedy search scores are unnormalized logits.
normalize_logits = True
else:
raise RuntimeError("Cannot postprocess the output of the model.")

transition_scores = cast(
torch.Tensor,
self.model.compute_transition_scores(
output_ids, # type: ignore
scores, # type: ignore
beam_indices, # type: ignore
normalize_logits=True,
),
)
transition_scores = _compute_transition_scores(output_ids, scores, beam_indices, normalize_logits)

if beam_indices is None:
beam_indices = torch.zeros_like(output_ids)
Expand Down Expand Up @@ -309,15 +324,26 @@ def _forward(self, model_inputs, **generate_kwargs):
start_index = 0
if self.model.config.decoder_start_token_id is not None:
start_index = 1
if generate_kwargs["output_attentions"] is True:
assert attentions is not None
# output_attentions can be unset or overridden per call, so rely on what generate actually returned.
if attentions is not None:

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Minor: F11. This guard misses the SDPA case: generate returns cross_attentions as a tuple of empty tuples, not None, so attentions[0][0] on the next line raises IndexError. Repro on stas/tiny-m2m_100: load with attn_implementation="sdpa", then call SilTranslationPipeline(..., output_attentions=True), and the call raises IndexError: tuple index out of range. Checking attentions and attentions[0] would fall back to empty alignments.

num_heads = attentions[0][0].shape[1]

# Truncate/Pad beam_indices to match output_ids length exact slice
target_seq_len = output_ids.shape[1] - start_index
sliced_beam_indices = beam_indices[:, start_index:]
Comment thread
claude[bot] marked this conversation as resolved.
Comment thread
claude[bot] marked this conversation as resolved.
if sliced_beam_indices.shape[1] > target_seq_len:
sliced_beam_indices = sliced_beam_indices[:, :target_seq_len]
elif sliced_beam_indices.shape[1] < target_seq_len:
sliced_beam_indices = torch.nn.functional.pad(
Comment thread
ddaspit marked this conversation as resolved.
sliced_beam_indices, (0, target_seq_len - sliced_beam_indices.shape[1])
)
Comment thread
claude[bot] marked this conversation as resolved.

indices = torch.stack(
(
torch.arange(output_ids.shape[1] - start_index, device=output_ids.device).expand(
in_b, n_sequences, -1
),
torch.reshape(beam_indices[:, start_index:] % num_beams, (in_b, n_sequences, -1)),
torch.reshape(sliced_beam_indices % num_beams, (in_b, n_sequences, -1)),
),
dim=3,
)
Expand Down Expand Up @@ -446,5 +472,43 @@ def torch_gather_nd(params: torch.Tensor, indices: torch.Tensor, batch_dim: int
return out.reshape(*index_shape, *tail_sizes)


def _compute_transition_scores(
Comment thread
ddaspit marked this conversation as resolved.
sequences: torch.Tensor,
scores: Tuple[torch.Tensor, ...],
output_beam_indices: Optional[torch.Tensor],
normalize_logits: bool,
) -> torch.Tensor:
"""
Compute the log probability of each generated token.

This is equivalent to PreTrainedModel.compute_transition_scores, but gathers one step at a time instead of
stacking the scores for every step into a single tensor, which requires several gigabytes of memory for a large
vocabulary.
"""
if output_beam_indices is None:
# Greedy search is equivalent to a beam search where the first (and only) beam is always selected.
beam_indices = torch.arange(scores[0].shape[0], device=sequences.device).view(-1, 1).expand(-1, len(scores))
else:
beam_indices = output_beam_indices

# Cut the beam indices to the longest beam length. Beams that finished early are masked out below.
beam_indices_mask = beam_indices < 0
max_beam_length = int((1 - beam_indices_mask.long()).sum(-1).max().item())
beam_indices_mask = beam_indices_mask[:, :max_beam_length]
beam_indices = beam_indices[:, :max_beam_length].masked_fill(beam_indices_mask, 0)

# The token generated at step i is at cut_idx + i in the sequence.
cut_idx = sequences.shape[-1] - max_beam_length
transition_scores = torch.zeros(sequences.shape[0], max_beam_length, device=sequences.device, dtype=scores[0].dtype)
for i in range(max_beam_length):
step_scores = scores[i]
if normalize_logits:
step_scores = torch.nn.functional.log_softmax(step_scores, dim=-1)
transition_scores[:, i] = step_scores[beam_indices[:, i], sequences[:, cut_idx + i]]

transition_scores[beam_indices_mask] = 0
return transition_scores


def _get_encoding_fast_tokens(encoding) -> List[str]:
return [token for (token, mask) in zip(encoding.tokens, encoding.special_tokens_mask) if not mask]
Loading
Loading