From 4f300c8632bd1ea75d50497f9557dddc5f5c592d Mon Sep 17 00:00:00 2001 From: "hemu.zp" Date: Sun, 16 Apr 2023 02:05:41 +0800 Subject: [PATCH] Fix generate for ModelForTextGeneration Separate the `generate` function, no longer use the implementation in the transformers library to avoid error due to transformers version upgrades. Link: https://code.alibaba-inc.com/Ali-MaaS/MaaS-lib/codereview/12340643 --- .../models/nlp/task_models/text_generation.py | 167 ++++++++++++++++-- tests/pipelines/test_text_generation.py | 4 +- 2 files changed, 154 insertions(+), 17 deletions(-) diff --git a/modelscope/models/nlp/task_models/text_generation.py b/modelscope/models/nlp/task_models/text_generation.py index e09a25ff..85d9066a 100644 --- a/modelscope/models/nlp/task_models/text_generation.py +++ b/modelscope/models/nlp/task_models/text_generation.py @@ -2,6 +2,7 @@ from typing import Any, Dict import numpy as np +import torch from transformers.modeling_utils import PreTrainedModel from modelscope.metainfo import TaskModels @@ -36,17 +37,17 @@ class ModelForTextGeneration(SingleBackboneTaskModelBase, PreTrainedModel): output_embeddings = self.head.get_output_embeddings() output_embeddings.weight = input_embeddings.weight - def forward(self, **input: Dict[str, Any]) -> Dict[str, np.ndarray]: + def forward(self, **inputs) -> Dict[str, np.ndarray]: # backbone do not need labels, only head need for loss compute - labels = input.pop(OutputKeys.LABELS, None) + labels = inputs.pop(OutputKeys.LABELS, None) - backbone_outputs = super().forward(input) + backbone_outputs = super().forward(inputs) hidden_states = backbone_outputs[0] logits = self.head.forward(hidden_states) loss = None if labels is not None: - input[OutputKeys.LABELS] = labels + inputs[OutputKeys.LABELS] = labels loss = self.compute_loss(logits, labels) return TextGenerationModelOutput(logits=logits, loss=loss) @@ -74,14 +75,150 @@ class ModelForTextGeneration(SingleBackboneTaskModelBase, PreTrainedModel): 'attention_mask': attention_mask, } - def generate(self, inputs, *args, **kwargs): - input_ids = inputs['input_ids'] if isinstance(inputs, Dict) else inputs - generate_output = super().generate(input_ids, *args, **kwargs) - if isinstance(generate_output, Dict): - return TokenGeneratorOutput( - sequences=generate_output.sequences, - scores=generate_output.scores, - attentions=generate_output.attentions, - hidden_states=generate_output.hidden_states) - else: - return TokenGeneratorOutput(sequences=generate_output) + def generate(self, inputs, temperature=1.0, **kwargs): + tokens = inputs['input_ids'] if isinstance(inputs, Dict) else inputs + top_k = kwargs.pop('top_k', + self.config.top_k if 'top_k' in self.config else 1) + top_p = kwargs.pop('top_p', + self.config.top_p if 'top_p' in self.config else 0.) + max_length = kwargs.pop('max_length', self.config.max_length) + + batch_size = tokens.size(0) + lengths = kwargs.pop( + 'prompt_length', + torch.tensor([tokens.size(1)], device=tokens.device)) + + min_prompt_length = lengths.min().item() + max_sequence_length = max_length + + # If the context is too big, this happens + if min_prompt_length >= max_sequence_length: + raise ValueError('context length too large') + + pad_length = max_sequence_length - tokens.size(1) + if pad_length > 0: + pads = torch.zeros( + batch_size, pad_length, device=tokens.device).long() + tokens = torch.cat((tokens, pads), dim=-1) + + # Added termination_id to support the case that we want to terminate the + # generation once that id is generated. + termination_id = self.config.eos_token_id + + # Whether we have reached a termination id. + is_generation_done = torch.zeros( + batch_size, dtype=torch.uint8, device=tokens.device) + + with torch.no_grad(): + for context_length in range(min_prompt_length, + max_sequence_length): + + # Pick the slice that we need to pass through the network. + tokens2use = tokens[:, :context_length] + + # logits will be meanigful only in the last pipeline stage. + logits = self(input_ids=tokens2use).logits + + # Sample. + last_token_logits = logits[:, -1, :] + new_sample = sample( + last_token_logits, + top_k=top_k, + top_p=top_p, + temperature=temperature, + vocab_size=self.head_cfg.vocab_size) + + # If a prompt length is smaller or equal th current context + # length, it means we have started generating tokens + started = lengths <= context_length + # Update the tokens. + tokens[started, context_length] = new_sample[started] + + done_token = (new_sample == termination_id).byte() & \ + started.byte() + + is_generation_done = is_generation_done | done_token + done = torch.all(is_generation_done) + + if done: + break + + tokens = tokens[:, :(context_length + 1)] + return TokenGeneratorOutput(sequences=tokens) + + +def sample(logits, top_k=0, top_p=0.0, temperature=1.0, vocab_size=None): + """ Sample and generate a token. + Note: logits has the dimension [b, v] where b is the batch size + and v is the vocabulary size. + If vocab_size is provided, we will make sure the sample that is + generated is in [0, vocab-size). This will avoid out of vocabulary + generations due to padding. + """ + + # Check logits for consistency. + assert logits.ndim == 2, 'expected the logits to be of [b, v] shape.' + + # Greedy is just simple argmax. + if top_k == 1: + assert top_p == 0.0, 'cannot set both greedy and top-p samplings.' + samples = torch.argmax(logits, dim=-1) + + # Top-k or top-p sampling. + else: + # Clone so we do not modify the inputs, + logits = logits.clone() + # Apply temperature in place. + if temperature != 1.0: + logits.div_(temperature) + + if top_k > 1: + top_p == 0.0 + assert top_k <= logits.size(1), 'top-k is larger than logit size.' + if vocab_size: + assert top_k < vocab_size, 'top-k is larger than vocab size.' + modify_logits_for_top_k_filtering(logits, top_k) + + elif top_p > 0.0: + assert top_p <= 1.0, 'top-p should be in (0, 1].' + modify_logits_for_top_p_filtering(logits, top_p) + + # After filtering, we need to recalculate the distribution. + probs = logits.softmax(dim=-1) + samples = torch.multinomial(probs, num_samples=1).view(-1) + + # If vocab size is provided, make sure the samples are in + # in the range [0, vocab-size). + if vocab_size: + samples = torch.clamp(samples, min=0, max=(vocab_size - 1)) + + return samples + + +def modify_logits_for_top_k_filtering(logits, top_k): + """Set the logits for none top-k values to -inf.""" + + filter_ = logits < torch.topk(logits, top_k)[0][..., -1, None] + logits.masked_fill_(filter_, float('-Inf')) + + +def modify_logits_for_top_p_filtering(logits, top_p): + """Set the logits for none top-p values to -inf.""" + + # First sort and calculate cumulative sum of probabilities. + sorted_logits, sorted_indices = torch.sort(logits, descending=True) + cumulative_probs = sorted_logits.softmax(dim=-1).cumsum(dim=-1) + + # Filteration based on the cumulative sum. + filter_ = cumulative_probs > top_p + # This shift by 1 is weird and I cannot justify it. This existed + # in the original implementation: + # https://github.com/ari-holtzman/degen/blob/master/gen.py + # and I guess it is needed so keeping it for now. + filter_[:, 1:] = filter_[:, :-1].clone() + # Make sure we at least have one token to select from. + filter_[..., 0] = 0 + + # Fill in the filtered part + filter_ = filter_.scatter(1, sorted_indices, filter_) + logits.masked_fill_(filter_, float('-Inf')) diff --git a/tests/pipelines/test_text_generation.py b/tests/pipelines/test_text_generation.py index a729d4da..998cbd18 100644 --- a/tests/pipelines/test_text_generation.py +++ b/tests/pipelines/test_text_generation.py @@ -67,7 +67,7 @@ class TextGenerationTest(unittest.TestCase, DemoCompatibilityCheck): self.run_pipeline_with_model_id(self.palm_model_id_zh_base, self.palm_input_zh) - @unittest.skipUnless(test_level() >= -1, 'skip test in current test level') + @unittest.skipUnless(test_level() >= 0, 'skip test in current test level') def test_palm_zh_base_with_model_name_with_args(self): self.run_pipeline_with_model_id( self.palm_model_id_zh_base, @@ -106,7 +106,7 @@ class TextGenerationTest(unittest.TestCase, DemoCompatibilityCheck): self.run_pipeline_with_model_id(self.gpt3_base_model_id, self.gpt3_input) - @unittest.skipUnless(test_level() >= -1, 'skip test in current test level') + @unittest.skipUnless(test_level() >= 0, 'skip test in current test level') def test_gpt_base_with_model_name_with_args(self): self.run_pipeline_with_model_id( self.gpt3_base_model_id,