-
-
Notifications
You must be signed in to change notification settings - Fork 2
Migrate to transformers 5 #363
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
a0b81b8
bfd801e
ed43f14
b6ad763
bbd2887
1665355
8149f7a
53d92d8
340028f
8ac226f
fcf9486
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -1 +1,2 @@ | ||
| *.ipynb eol=lf | ||
| *.sh eol=lf |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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 | ||
|
|
@@ -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) | ||
|
|
@@ -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 | ||
|
|
@@ -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}'") | ||
|
|
||
|
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 | ||
|
|
@@ -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( | ||
|
claude[bot] marked this conversation as resolved.
|
||
| model=self._model, | ||
| tokenizer=self._tokenizer, | ||
| batch_size=self._batch_size, | ||
|
|
@@ -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: | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. FYI: F12. |
||
| # 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, | ||
|
|
@@ -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 | ||
|
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 | ||
|
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) | ||
|
|
@@ -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: | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Minor: F11. This guard misses the SDPA case: |
||
| 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:] | ||
|
claude[bot] marked this conversation as resolved.
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( | ||
|
ddaspit marked this conversation as resolved.
|
||
| sliced_beam_indices, (0, target_seq_len - sliced_beam_indices.shape[1]) | ||
| ) | ||
|
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, | ||
| ) | ||
|
|
@@ -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( | ||
|
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] | ||
Uh oh!
There was an error while loading. Please reload this page.