clef / code /models /common /generation_utils.py
tt-hous's picture
Add files using upload-large-folder tool
be3ecc8 verified
Raw History Blame Contribute Delete
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)