mirror of
https://github.com/modelscope/modelscope.git
synced 2026-09-01 19:49:03 +02:00
Support run text generation pipeline with args
Link: https://code.alibaba-inc.com/Ali-MaaS/MaaS-lib/codereview/11937122
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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, {}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user