mirror of
https://github.com/modelscope/modelscope.git
synced 2026-09-01 19:49:03 +02:00
llama2 support chat (#404)
* support chat * update llama2 chat testcase * add gen kwargs and devices * update unittest and support max_length in multi-turn dialogue
This commit is contained in:
@@ -17,7 +17,7 @@
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
from typing import List, Optional, Tuple, Union
|
||||
from typing import Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -27,11 +27,48 @@ from torch.nn import CrossEntropyLoss
|
||||
from transformers.modeling_outputs import CausalLMOutputWithPast
|
||||
|
||||
from modelscope.metainfo import Models
|
||||
from modelscope.outputs import OutputKeys
|
||||
from modelscope.utils.constant import Tasks
|
||||
from ... import MODELS
|
||||
from .backbone import Llama2Model, LlamaPreTrainedModel
|
||||
|
||||
|
||||
def get_chat_prompt(system: str, text: str, history: List[Tuple[str, str]],
|
||||
max_length: int, tokenizer):
|
||||
system_prompt = f'<s>[INST] <<SYS>>\n{system}\n<</SYS>>\n\n'
|
||||
system_ids = tokenizer(system_prompt, return_tensors='pt').input_ids
|
||||
|
||||
text_prompt = f'{text.strip()} [/INST]'
|
||||
text_ids = tokenizer(text_prompt, return_tensors='pt').input_ids
|
||||
|
||||
prompt_length = system_ids.shape[-1] + text_ids.shape[-1]
|
||||
if prompt_length > max_length:
|
||||
raise RuntimeError(
|
||||
f'prepend prompt length {prompt_length} is bigger than max_length {max_length}'
|
||||
)
|
||||
|
||||
history_prompt = ''
|
||||
history_ids_list = []
|
||||
# traverse history in reverse order
|
||||
for user, bot in history[::-1]:
|
||||
assert isinstance(user, str)
|
||||
assert isinstance(bot, str)
|
||||
round_prompt = f'{user.strip()} [/INST] {bot.strip()} </s><s>[INST] '
|
||||
round_ids = tokenizer(round_prompt, return_tensors='pt').input_ids
|
||||
if prompt_length + round_ids.shape[-1] > max_length:
|
||||
# excess history should not be appended to the prompt
|
||||
break
|
||||
else:
|
||||
history_prompt = round_prompt + history_prompt
|
||||
history_ids_list = [round_ids] + history_ids_list
|
||||
prompt_length += round_ids.shape[-1]
|
||||
|
||||
prompt_list = [system_prompt, history_prompt, text_prompt]
|
||||
prompt_ids_list = [system_ids] + history_ids_list + [text_ids]
|
||||
|
||||
return ''.join(prompt_list), torch.cat(prompt_ids_list, dim=1)
|
||||
|
||||
|
||||
# This file is mainly copied from the llama code of transformers
|
||||
@MODELS.register_module(Tasks.text_generation, module_name=Models.llama2)
|
||||
class Llama2ForTextGeneration(LlamaPreTrainedModel):
|
||||
@@ -186,3 +223,46 @@ class Llama2ForTextGeneration(LlamaPreTrainedModel):
|
||||
past_state.index_select(0, beam_idx.to(past_state.device))
|
||||
for past_state in layer_past), )
|
||||
return reordered_past
|
||||
|
||||
def chat(self, input: Dict, tokenizer) -> Dict:
|
||||
import copy
|
||||
gen_kwargs = copy.copy(input)
|
||||
if 'text' not in input:
|
||||
text: str = ''
|
||||
else:
|
||||
text: str = input['text']
|
||||
gen_kwargs.pop('text')
|
||||
|
||||
if 'system' not in input:
|
||||
system: str = ''
|
||||
else:
|
||||
system: str = input['system']
|
||||
gen_kwargs.pop('system')
|
||||
|
||||
if 'history' not in input:
|
||||
history = []
|
||||
else:
|
||||
history: List[Tuple] = copy.copy(input['history'])
|
||||
gen_kwargs.pop('history')
|
||||
|
||||
if 'max_length' not in gen_kwargs:
|
||||
gen_kwargs['max_length'] = 4096
|
||||
|
||||
prompt, prompt_ids = get_chat_prompt(
|
||||
system=system,
|
||||
text=text,
|
||||
history=history,
|
||||
max_length=gen_kwargs['max_length'],
|
||||
tokenizer=tokenizer)
|
||||
input_ids = prompt_ids.to(self.device)
|
||||
generate_ids = self.generate(input_ids, **gen_kwargs)
|
||||
# remove input tokens
|
||||
generate_ids = generate_ids[:, input_ids.shape[1]:]
|
||||
response = tokenizer.batch_decode(
|
||||
generate_ids,
|
||||
skip_special_tokens=True,
|
||||
clean_up_tokenization_spaces=False)[0]
|
||||
response = response.strip()
|
||||
history.append((text, response))
|
||||
|
||||
return {OutputKeys.RESPONSE: response, OutputKeys.HISTORY: history}
|
||||
|
||||
59
tests/models/test_llama2.py
Normal file
59
tests/models/test_llama2.py
Normal file
@@ -0,0 +1,59 @@
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
|
||||
from modelscope import Model, snapshot_download
|
||||
from modelscope.models.nlp.llama2 import Llama2Tokenizer
|
||||
from modelscope.utils.test_utils import test_level
|
||||
|
||||
|
||||
class Llama2Test(unittest.TestCase):
|
||||
|
||||
def setUp(self) -> None:
|
||||
self.model_name = 'modelscope/Llama-2-7b-chat-ms'
|
||||
self.system = 'you are a helpful assistant!'
|
||||
self.text_first_round = 'hello'
|
||||
self.text_second_round = 'do you know peking university?'
|
||||
self.text_third_round = 'where is it?'
|
||||
|
||||
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
|
||||
def test_chat(self):
|
||||
model_dir = snapshot_download(
|
||||
self.model_name, ignore_file_pattern=[r'\w+\.safetensors'])
|
||||
model = Model.from_pretrained(
|
||||
model_dir, device_map='auto', torch_dtype=torch.float16)
|
||||
tokenizer = Llama2Tokenizer.from_pretrained(model_dir)
|
||||
|
||||
inputs = {
|
||||
'text': self.text_first_round,
|
||||
'history': [],
|
||||
'system': self.system
|
||||
}
|
||||
result = model.chat(input=inputs, tokenizer=tokenizer)
|
||||
self.assertIsInstance(result['history'], list)
|
||||
self.assertEqual(len(result['history']), 1)
|
||||
self.assertEqual(result['history'][0][0], self.text_first_round)
|
||||
|
||||
inputs = {
|
||||
'text': self.text_second_round,
|
||||
'history': result['history'],
|
||||
'system': self.system
|
||||
}
|
||||
result = model.chat(input=inputs, tokenizer=tokenizer)
|
||||
self.assertIsInstance(result['history'], list)
|
||||
self.assertEqual(len(result['history']), 2)
|
||||
self.assertEqual(result['history'][1][0], self.text_second_round)
|
||||
|
||||
inputs = {
|
||||
'text': self.text_third_round,
|
||||
'history': result['history'],
|
||||
'system': self.system
|
||||
}
|
||||
result = model.chat(input=inputs, tokenizer=tokenizer)
|
||||
self.assertIsInstance(result['history'], list)
|
||||
self.assertEqual(len(result['history']), 3)
|
||||
self.assertEqual(result['history'][2][0], self.text_third_round)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user