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:
lylalala
2023-07-24 15:38:01 +08:00
committed by GitHub
parent 9dedac31fd
commit f805d86aed
2 changed files with 140 additions and 1 deletions

View File

@@ -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}

View 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()