Download code/models/common/generation_utils.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 11.4 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/generation_utils.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/generation_utils.py
-
curl -L -o generation_utils.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/generation_utils.py
11.4 kB
| # SPDX-FileCopyrightText: © 2023 Tenstorrent USA, Inc. | |
| # SPDX-License-Identifier: Apache-2.0 | |
| import torch | |
| from loguru import logger | |
| from transformers.generation.configuration_utils import GenerationConfig | |
| from transformers.generation.logits_process import ( # ForceTokensLogitsProcessor, | |
| EncoderNoRepeatNGramLogitsProcessor, | |
| EncoderRepetitionPenaltyLogitsProcessor, | |
| ExponentialDecayLengthPenalty, | |
| ForcedBOSTokenLogitsProcessor, | |
| ForcedEOSTokenLogitsProcessor, | |
| InfNanRemoveLogitsProcessor, | |
| LogitNormalization, | |
| LogitsProcessorList, | |
| MinLengthLogitsProcessor, | |
| MinNewTokensLengthLogitsProcessor, | |
| NoBadWordsLogitsProcessor, | |
| NoRepeatNGramLogitsProcessor, | |
| PrefixConstrainedLogitsProcessor, | |
| RepetitionPenaltyLogitsProcessor, | |
| SuppressTokensAtBeginLogitsProcessor, | |
| SuppressTokensLogitsProcessor, | |
| ) | |
| # HammingDiversityLogitsProcessor (diverse beam search) was removed in | |
| # transformers 5.x with no replacement. Import it optionally so this module | |
| # still loads; it's only used when diversity_penalty > 0, which TT generation | |
| # paths don't exercise. | |
| try: | |
| from transformers.generation.logits_process import HammingDiversityLogitsProcessor | |
| except ImportError: # transformers >= 5.x | |
| HammingDiversityLogitsProcessor = None | |
| def _merge_criteria_processor_list( | |
| default_list, # Union[LogitsProcessorList, StoppingCriteriaList], | |
| custom_list, # Union[LogitsProcessorList, StoppingCriteriaList], | |
| ): # -> Union[LogitsProcessorList, StoppingCriteriaList]: | |
| if len(custom_list) == 0: | |
| return default_list | |
| for default in default_list: | |
| for custom in custom_list: | |
| if type(custom) is type(default): | |
| object_type = "stopping criteria" if isinstance(custom, StoppingCriteria) else "logits processor" | |
| raise ValueError( | |
| f"A custom {object_type} of type {type(custom)} with values {custom} has been passed to" | |
| f" `generate`, but it has already been created with the values {default}. {default} has been" | |
| " created by passing the corresponding arguments to generate or by the model's config default" | |
| f" values. If you just want to change the default values of {object_type} consider passing" | |
| f" them as arguments to `generate` instead of using a custom {object_type}." | |
| ) | |
| default_list.extend(custom_list) | |
| return default_list | |
| def _get_logits_processor( | |
| generation_config: GenerationConfig, | |
| input_ids_seq_length: int, | |
| encoder_input_ids, # torch.LongTensor | |
| prefix_allowed_tokens_fn, # Callable[[int, torch.Tensor], List[int]], | |
| logits_processor, # Optional[LogitsProcessorList] | |
| ): # -> LogitsProcessorList: | |
| """ | |
| This class returns a [`LogitsProcessorList`] list object that contains all relevant [`LogitsProcessor`] | |
| instances used to modify the scores of the language model head. | |
| """ | |
| # instantiate processors list | |
| processors = LogitsProcessorList() | |
| # the following idea is largely copied from this PR: https://github.com/huggingface/transformers/pull/5420/files | |
| # all samplers can be found in `generation_utils_samplers.py` | |
| if generation_config.diversity_penalty is not None and generation_config.diversity_penalty > 0.0: | |
| if HammingDiversityLogitsProcessor is None: | |
| raise NotImplementedError( | |
| "diversity_penalty > 0 (diverse beam search) requires HammingDiversityLogitsProcessor, " | |
| "which was removed in transformers 5.x." | |
| ) | |
| processors.append( | |
| HammingDiversityLogitsProcessor( | |
| diversity_penalty=generation_config.diversity_penalty, | |
| num_beams=generation_config.num_beams, | |
| num_beam_groups=generation_config.num_beam_groups, | |
| ) | |
| ) | |
| if generation_config.encoder_repetition_penalty is not None and generation_config.encoder_repetition_penalty != 1.0: | |
| processors.append( | |
| EncoderRepetitionPenaltyLogitsProcessor( | |
| penalty=generation_config.encoder_repetition_penalty, | |
| encoder_input_ids=encoder_input_ids, | |
| ) | |
| ) | |
| if generation_config.repetition_penalty is not None and generation_config.repetition_penalty != 1.0: | |
| processors.append(RepetitionPenaltyLogitsProcessor(penalty=generation_config.repetition_penalty)) | |
| if generation_config.no_repeat_ngram_size is not None and generation_config.no_repeat_ngram_size > 0: | |
| processors.append(NoRepeatNGramLogitsProcessor(generation_config.no_repeat_ngram_size)) | |
| if ( | |
| generation_config.encoder_no_repeat_ngram_size is not None | |
| and generation_config.encoder_no_repeat_ngram_size > 0 | |
| ): | |
| if len(encoder_input_ids.shape) == 2: | |
| processors.append( | |
| EncoderNoRepeatNGramLogitsProcessor(generation_config.encoder_no_repeat_ngram_size, encoder_input_ids) | |
| ) | |
| else: | |
| raise ValueError("It's impossible to use `encoder_no_repeat_ngram_size` with decoder-only architecture") | |
| if generation_config.bad_words_ids is not None: | |
| processors.append(NoBadWordsLogitsProcessor(generation_config.bad_words_ids, generation_config.eos_token_id)) | |
| if ( | |
| generation_config.min_length is not None | |
| and generation_config.eos_token_id is not None | |
| and generation_config.min_length > 0 | |
| ): | |
| processors.append(MinLengthLogitsProcessor(generation_config.min_length, generation_config.eos_token_id)) | |
| if ( | |
| generation_config.min_new_tokens is not None | |
| and generation_config.eos_token_id is not None | |
| and generation_config.min_new_tokens > 0 | |
| ): | |
| processors.append( | |
| MinNewTokensLengthLogitsProcessor( | |
| input_ids_seq_length, | |
| generation_config.min_new_tokens, | |
| generation_config.eos_token_id, | |
| ) | |
| ) | |
| if prefix_allowed_tokens_fn is not None: | |
| processors.append( | |
| PrefixConstrainedLogitsProcessor( | |
| prefix_allowed_tokens_fn, | |
| generation_config.num_beams // generation_config.num_beam_groups, | |
| ) | |
| ) | |
| if generation_config.forced_bos_token_id is not None: | |
| processors.append(ForcedBOSTokenLogitsProcessor(generation_config.forced_bos_token_id)) | |
| if generation_config.forced_eos_token_id is not None: | |
| processors.append( | |
| ForcedEOSTokenLogitsProcessor(generation_config.max_length, generation_config.forced_eos_token_id) | |
| ) | |
| if generation_config.remove_invalid_values is True: | |
| processors.append(InfNanRemoveLogitsProcessor()) | |
| if generation_config.exponential_decay_length_penalty is not None: | |
| processors.append( | |
| ExponentialDecayLengthPenalty( | |
| generation_config.exponential_decay_length_penalty, | |
| generation_config.eos_token_id, | |
| input_ids_seq_length, | |
| ) | |
| ) | |
| if generation_config.suppress_tokens is not None: | |
| processors.append(SuppressTokensLogitsProcessor(generation_config.suppress_tokens)) | |
| if generation_config.begin_suppress_tokens is not None: | |
| begin_index = input_ids_seq_length | |
| begin_index = ( | |
| begin_index | |
| if (input_ids_seq_length > 1 or generation_config.forced_bos_token_id is None) | |
| else begin_index + 1 | |
| ) | |
| processors.append(SuppressTokensAtBeginLogitsProcessor(generation_config.begin_suppress_tokens, begin_index)) | |
| processors = _merge_criteria_processor_list(processors, logits_processor) | |
| # `LogitNormalization` should always be the last logit processor, when present | |
| if generation_config.renormalize_logits is True: | |
| processors.append(LogitNormalization()) | |
| return processors | |
| def get_logits_processor(input_ids, config): | |
| generation_config = GenerationConfig.from_model_config(config) | |
| input_ids_seq_length = input_ids.shape[-1] | |
| logits_processor = _get_logits_processor( | |
| generation_config=generation_config, | |
| input_ids_seq_length=input_ids_seq_length, | |
| encoder_input_ids=input_ids, | |
| prefix_allowed_tokens_fn=None, | |
| logits_processor=LogitsProcessorList(), | |
| ) | |
| return logits_processor | |
| def pad_input_32(tensor, value): | |
| len = tensor.shape[1] | |
| if len % 32 == 0: | |
| return tensor | |
| padded_len = ((len // 32) + 1) * 32 | |
| pad_tensor = (value * torch.ones(tensor.shape[0], padded_len - len)).to(torch.long) | |
| tensor = torch.cat([tensor, pad_tensor], dim=1) | |
| return tensor | |
| def run_generate( | |
| input_sentance, | |
| tokenizer, | |
| tt_model_constructor, | |
| device, | |
| run_tt_model=True, | |
| log=True, | |
| comp_pcc=None, | |
| ): | |
| tt_model, hf_reference_model = tt_model_constructor(device) | |
| # Prepare input | |
| tokenized = tokenizer(input_sentance, return_tensors="pt") # Batch size 1 | |
| input_ids = pad_input_32(tokenized.input_ids, hf_reference_model.generation_config.pad_token_id) | |
| attention_mask = pad_input_32(tokenized.attention_mask, 0) | |
| if log: | |
| logger.debug(f"input_ids {input_ids.shape} {input_ids}") | |
| logger.debug(f"attention_mask {attention_mask.shape} {attention_mask}") | |
| logits_processor = get_logits_processor(input_ids, hf_reference_model.config) | |
| decoder_start_values = hf_reference_model.generation_config.pad_token_id * torch.ones(1, 32).to(torch.long) | |
| decoder_input_ids = hf_reference_model.generation_config.pad_token_id * torch.ones(1, 64).to(torch.long) | |
| if log: | |
| logger.debug(f"decoder_input_ids {decoder_input_ids}") | |
| encoder_outputs = None | |
| use_cache = False | |
| for i in range(64): | |
| # PyTorch forward pass | |
| pt_out = hf_reference_model( | |
| input_ids=input_ids, | |
| decoder_input_ids=decoder_input_ids, | |
| attention_mask=attention_mask, | |
| ) | |
| if run_tt_model: | |
| tt_out = tt_model( | |
| input_ids=input_ids, | |
| decoder_input_ids=decoder_input_ids, | |
| attention_mask=attention_mask, | |
| encoder_outputs=encoder_outputs, | |
| return_dict=True, | |
| use_cache=use_cache, | |
| ) | |
| encoder_outputs = tt_out.encoder_outputs | |
| next_token_logits = tt_out.logits | |
| if comp_pcc is not None: | |
| does_pass, pcc_message = comp_pcc(pt_out.logits, tt_out.logits, 0.98) | |
| if log: | |
| logger.info(pcc_message) | |
| else: | |
| next_token_logits = pt_out.logits | |
| # pre-process distribution | |
| next_tokens_scores = logits_processor(input_ids, next_token_logits) | |
| # argmax | |
| next_tokens = torch.argmax(next_tokens_scores, dim=-1) | |
| if log: | |
| logger.debug(f"next_tokens {next_tokens}") | |
| if next_tokens[0][i] == hf_reference_model.generation_config.eos_token_id: | |
| break | |
| # We need to expand decoder_input_ids | |
| if (i + 1) % 32 == 0: | |
| decoder_input_ids = torch.cat([decoder_input_ids, decoder_start_values], dim=1) | |
| decoder_input_ids[0][i + 1] = next_tokens[0][i] | |
| if log: | |
| logger.debug(f"decoder_input_ids {decoder_input_ids[0]}") | |
| return tokenizer.decode(decoder_input_ids[0], skip_special_tokens=True) | |