Support run text generation pipeline with args

Link: https://code.alibaba-inc.com/Ali-MaaS/MaaS-lib/codereview/11937122
This commit is contained in:
hemu.zp
2023-03-10 09:48:10 +08:00
committed by wenmeng.zwm
parent e02a260c93
commit 429cfee826
9 changed files with 96 additions and 43 deletions

View File

@@ -354,6 +354,9 @@ class GPT3Model(PreTrainedModel):
return model
def generate(self, tokens, temperature=1.0, **kwargs):
top_k = kwargs.pop('top_k', self.config.top_k)
top_p = kwargs.pop('top_p', self.config.top_p)
max_length = kwargs.pop('max_length', tokens.size(1) + 100)
batch_size = tokens.size(0)
lengths = kwargs.pop(
@@ -361,13 +364,18 @@ class GPT3Model(PreTrainedModel):
torch.tensor([tokens.size(1)], device=tokens.device))
min_prompt_length = lengths.min().item()
max_sequence_length = tokens.size(1)
max_sequence_length = min(max_sequence_length,
max_sequence_length = min(max_length,
self.config.max_position_embeddings)
# If the context is too big, this happens
if min_prompt_length >= max_sequence_length:
raise ValueError('context length + tokens_to_generate too large')
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.
@@ -391,8 +399,8 @@ class GPT3Model(PreTrainedModel):
last_token_logits = logits[:, -1, :]
new_sample = sample(
last_token_logits,
top_k=self.config.top_k,
top_p=self.config.top_p,
top_k=top_k,
top_p=top_p,
temperature=temperature,
vocab_size=self.config.vocab_size)

View File

@@ -974,8 +974,8 @@ class DistributedGPT3(TorchModel):
self.dist_model = model
tensor_ws = mpu.get_tensor_model_parallel_world_size()
ckpt_ws = get_args().get('checkpoint_tensor_model_parallel_size',
tensor_ws)
ckpt_ws = get_args().get('checkpoint_tensor_model_parallel_size', None)
ckpt_ws = tensor_ws if ckpt_ws is None else ckpt_ws
ckpt_rank = mpu.get_tensor_model_parallel_rank() * ckpt_ws // tensor_ws
load_model = pre_load(ckpt_rank, model_dir, tag=path_load_tag)
load_model = split_state_dict(load_model, model, tensor_ws // ckpt_ws)
@@ -1032,24 +1032,32 @@ class DistributedGPT3(TorchModel):
stop_on_double_eol=False,
stop_on_eol=False,
**kwargs):
top_k = kwargs.pop('top_k', self.config.top_k)
top_p = kwargs.pop('top_p', self.config.top_p)
temperature = kwargs.pop('temperature', self.config.temperature)
max_length = kwargs.pop(
'max_length',
tokens.size(1) + self.config.tokens_to_generate)
batch_size = tokens.size(0)
lengths = prompts_len
if lengths is None:
lengths = torch.tensor([tokens.size(1)], device=tokens.device)
pads = torch.ones(
batch_size, self.config.tokens_to_generate,
device=tokens.device).long() * self.config.eod_id
tokens = torch.cat((tokens, pads), dim=-1)
min_prompt_length = lengths.min().item()
max_sequence_length = tokens.size(1)
max_sequence_length = min(max_sequence_length,
max_sequence_length = min(max_length,
self.config.max_position_embeddings)
# If the context is too big, this happens
if min_prompt_length >= max_sequence_length:
raise ValueError('context length + tokens_to_generate 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)
# Initialize inference parameters.
self.inference_params = InferenceParams(batch_size,
max_sequence_length)
@@ -1084,9 +1092,9 @@ class DistributedGPT3(TorchModel):
last_token_logits = logits[:, -1, :]
new_sample = sample(
last_token_logits,
top_k=kwargs.pop('top_k', self.config.top_k),
top_p=kwargs.pop('top_p', self.config.top_p),
temperature=kwargs.pop('temperature', self.config.temperature),
top_k=top_k,
top_p=top_p,
temperature=temperature,
vocab_size=self.config.vocab_size)
# If a prompt length is smaller or equal th current context

View File

@@ -52,13 +52,14 @@ class GPT3ForTextGeneration(TorchModel):
"""
return self.model(**input)
def generate(self, inputs: Dict[str, Tensor]) -> Dict[str, Tensor]:
def generate(self, inputs: Dict[str, Tensor],
**kwargs) -> Dict[str, Tensor]:
if not isinstance(self.model, GPT3Model):
return self.model.generate(**inputs)
return self.model.generate(**inputs, **kwargs)
tokens = inputs['input_ids']
lengths = self._get_length(inputs['attention_mask'])
return self.model.generate(tokens, prompt_length=lengths)
return self.model.generate(tokens, prompt_length=lengths, **kwargs)
@staticmethod
def _get_length(attention_mask: torch.Tensor) -> Tensor:

View File

@@ -779,8 +779,6 @@ class Translator(object):
self.end_token = self.symbols['EOS']
self.alpha = self.args.alpha
self.beam_size = self.args.beam_size
self.min_length = self.args.min_length
self.max_length = self.args.max_length
def from_batch(self, translation_batch):
batch = translation_batch['batch']
@@ -1065,8 +1063,7 @@ class Translator(object):
"""
self.model.eval()
with torch.no_grad():
return self._fast_translate_batch(
batch, self.max_length, min_length=self.min_length)
return self._fast_translate_batch(batch)
def _tile(self, x, count, dim=0):
perm = list(range(len(x.size())))
@@ -1121,13 +1118,13 @@ class Translator(object):
logits[indices_to_remove] = filter_value
return logits
def _fast_translate_batch(self,
batch: 'Batch',
max_length: int,
min_length: int = 0):
def _fast_translate_batch(self, batch: 'Batch'):
# TODO: faster code path for beam_size == 1.
# TODO: support these blacklisted features.
max_length = self.args.max_length
min_length = self.args.min_length
beam_size = self.beam_size
batch_size = batch.batch_size
src = batch.src
@@ -1366,7 +1363,10 @@ class PalmForTextGeneration(PalmPreTrainedModel):
logits=output[0],
)
def generate(self, input: Dict[str, Tensor]) -> TokenGeneratorOutput:
def generate(self, input: Dict[str, Tensor],
**kwargs) -> TokenGeneratorOutput:
for k, v in kwargs.items():
setattr(self.generator.args, k, v)
outputs = self.generator(**input)
preds = outputs['predictions']
return TokenGeneratorOutput(sequences=[pred[0] for pred in preds])

View File

@@ -2,6 +2,7 @@
import os
import os.path as osp
import random
from abc import ABC, abstractmethod
from functools import partial
from multiprocessing import Pool
@@ -436,15 +437,20 @@ class DistributedPipeline(Pipeline):
ranks = list(range(self.world_size))
self.model_pool = Pool(self.world_size)
master_ip = '127.0.0.1' if 'master_ip' not in kwargs else kwargs[
'master_ip']
os.environ['MASTER_ADDR'] = master_ip
master_port = '29500' if 'master_port' not in kwargs else kwargs[
'master_port']
if 'master_ip' not in kwargs:
kwargs['master_ip'] = '127.0.0.1'
master_port = int(kwargs['master_port']
) if 'master_port' in kwargs else random.randint(
29500, 39500)
from modelscope.utils.torch_utils import _find_free_port, _is_free_port
if not _is_free_port(int(master_port)):
master_port = str(_find_free_port())
os.environ['MASTER_PORT'] = master_port
if not _is_free_port(master_port):
master_port = _find_free_port()
kwargs['master_port'] = str(master_port)
# TODO: Pass ip and port to megatron_util for initialization
os.environ['MASTER_ADDR'] = kwargs['master_ip']
os.environ['MASTER_PORT'] = kwargs['master_port']
self.model_pool.map(
partial(
self.__class__._instantiate_one,

View File

@@ -43,7 +43,7 @@ class DistributedGPT3Pipeline(DistributedPipeline):
def _forward_one(cls, inputs: Dict[str, Any]) -> Dict[str, Any]:
tokens = inputs['inputs']['input_ids'].cuda(
torch.cuda.current_device())
return cls.model.generate(tokens)
return cls.model.generate(tokens, **inputs['forward_params'])
def postprocess(self, inputs: Dict[str, Any],
**postprocess_params) -> Dict[str, str]:
@@ -61,3 +61,6 @@ class DistributedGPT3Pipeline(DistributedPipeline):
self.preprocessor.tokenizer.detokenize(
inputs.sequences[0].tolist())
}
def _sanitize_parameters(self, **pipeline_parameters):
return {}, pipeline_parameters, {}

View File

@@ -107,7 +107,7 @@ class TextGenerationTransformersPreprocessor(TextGenerationPreprocessorBase):
mode: str = ModeKeys.INFERENCE,
src_txt='src_txt',
tgt_txt='tgt_txt',
max_length: int = None,
sequence_length: int = None,
use_fast: bool = None,
keep_original_columns=None,
**kwargs):
@@ -118,7 +118,7 @@ class TextGenerationTransformersPreprocessor(TextGenerationPreprocessorBase):
mode: The mode for the preprocessor.
src_txt: The key of the source sentence.
tgt_txt: The key of the generated sentence.
max_length: The max sequence length which the model supported,
sequence_length: The max sequence length which the model supported,
will be passed into tokenizer as the 'max_length' param.
use_fast: Whether to use the fast tokenizer or not.
**kwargs: Extra args input into the tokenizer's __call__ method.
@@ -130,10 +130,10 @@ class TextGenerationTransformersPreprocessor(TextGenerationPreprocessorBase):
kwargs['padding'] = kwargs.get('padding', 'max_length')
kwargs['return_token_type_ids'] = kwargs.get('return_token_type_ids',
False)
kwargs[
'max_length'] = max_length if max_length is not None else kwargs.get(
'sequence_length', 128)
kwargs.pop('sequence_length', None)
# sequence_length > max_length
kwargs['max_length'] = sequence_length if sequence_length is not None \
else kwargs.get('max_length', 128)
self.src_length = kwargs['max_length']
self.tgt_length = kwargs.pop('target_max_length', kwargs['max_length'])
model_type = None

View File

@@ -27,6 +27,11 @@ class TextGPT3GenerationTest(unittest.TestCase):
pipe = pipeline(Tasks.text_generation, model=self.model_id_2_7B)
print(pipe(self.input))
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
def test_gpt3_1_3B_with_args(self):
pipe = pipeline(Tasks.text_generation, model=self.model_id_1_3B)
print(pipe(self.input, top_p=0.9, temperature=0.9, max_length=32))
@unittest.skip('distributed gpt3 13B, skipped')
def test_gpt3_13B(self):
""" The model can be downloaded from the link on

View File

@@ -67,6 +67,17 @@ 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')
def test_palm_zh_base_with_model_name_with_args(self):
self.run_pipeline_with_model_id(
self.palm_model_id_zh_base,
self.palm_input_zh,
run_kwargs={
'top_p': 0.9,
'temperature': 0.9,
'max_length': 64
})
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
def test_palm_zh_base_with_model_name_batch(self):
self.run_pipeline_with_model_id(
@@ -95,6 +106,17 @@ 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')
def test_gpt_base_with_model_name_with_args(self):
self.run_pipeline_with_model_id(
self.gpt3_base_model_id,
self.gpt3_input,
run_kwargs={
'top_p': 0.9,
'temperature': 0.9,
'max_length': 64
})
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
def test_gpt_base_with_model_name_batch(self):
self.run_pipeline_with_model_id(