已合并
【HunyuanLargeMoE】part of data-preprocess #2248
zhoubeirong创建于 2025年2月19日
【HunyuanLargeMoE】part of data-preprocess #2248
已合并
zhoubeirong创建于 2025年2月19日
从refs/pull/2248/head合入到master
共 6 个文件变更+187-7
@@ -562,5 +562,19 @@
562 "stop_words": [562 "stop_words": [
563 "<|im_end|>"563 "<|im_end|>"
564 ]564 ]
565+ },
566+ {
567+ "name": "hunyuan",
568+ "format_user": {
569+ "slots": [
570+ "{{content}}"
571+ ]
572+ },
573+ "format_assistant": {
574+ "slots": [
575+ "{{content}}"
576+ ]
577+ },
578+ "default_system": "You are a helpful assistant."
565 }579 }
566]580]
@@ -398,6 +398,76 @@ class SharegptStyleInstructionHandler(LlamaFactoryInstructionHandler):
398 super().__init__(args, raw_datasets, tokenizer, splitter)398 super().__init__(args, raw_datasets, tokenizer, splitter)
399 399 
400 400 
401+class HunyuanInstructionHandler(BaseDatasetHandler):
402+ """
403+ Handle HunyuanLarge supported dataset format
404+ a HunyuanLarge instruction dataset handler
405+ """
406+ 
407+ def __init__(self, args, raw_datasets, tokenizer, splitter):
408+ super().__init__(args, raw_datasets, tokenizer, splitter)
409+ self.prompter = None
410+ self.train_on_inputs = False
411+ self.args.json_keys = ["input_ids", "attention_mask", "labels"]
412+ # use 'packed' string to mark that this is a packed dataset
413+ self.args.output_prefix = self.args.output_prefix + "_packed"
414+ self.ignored_label = -100
415+ self.is_multi_turn = True
416+ 
417+ self.hunyuanlarge_template = get_model_template(args.prompt_type.strip(), args.prompt_type_path.strip())
418+ 
419+ def _format_msg(self, sample):
420+ return sample
421+ 
422+ def apply_chat_template(self, conversations):
423+ return self.tokenizer.tokenizer.apply_chat_template(conversations)
424+ 
425+ def _tokenize_prompt(
426+ self,
427+ example,
428+ template,
429+ tokenizer,
430+ ) -> Dict[str, List[List[int]]]:
431+ model_inputs = {"input_ids": [], "attention_mask": [], "labels": []}
432+ input_ids, labels = [], []
433+ messages = example["prompt"] + example["response"]
434+ 
435+ message_tokens = torch.tensor(self.tokenizer.tokenizer.apply_chat_template(messages))
436+ 
437+ IGNORE_INDEX = -100
438+ 
439+ extra_0_token_id = self.tokenizer.tokenizer.convert_tokens_to_ids('<|extra_0|>')
440+ eos_token_id = self.tokenizer.tokenizer.convert_tokens_to_ids('<|eos|>')
441+ loss_token_begins = (message_tokens == extra_0_token_id).nonzero(as_tuple=True)[0].tolist()
442+ loss_token_ends = (message_tokens == eos_token_id).nonzero(as_tuple=True)[0].tolist()
443+ message_labels = torch.tensor([IGNORE_INDEX] * message_tokens.shape[0])
444+ for begin_idx, end_idx in zip(loss_token_begins, loss_token_ends):
445+ message_labels[begin_idx:end_idx + 1] = message_tokens[begin_idx:end_idx + 1]
446+ input_ids = message_tokens.to(torch.long)
447+ labels = message_labels.to(torch.long)
448+ 
449+ input_ids = input_ids[:self.max_seq_len]
450+ labels = labels[:self.max_seq_len]
451+ attention_mask = [1 if val != self.tokenizer.tokenizer.pad_id else 0 for val in input_ids]
452+ model_inputs["input_ids"] = input_ids
453+ model_inputs["attention_mask"] = torch.tensor(attention_mask, dtype=torch.bool)
454+ model_inputs["labels"] = labels
455+ 
456+ return model_inputs
457+ 
458+ def _filter(self, sample):
459+ messages = self._format_msg(sample)
460+ tokenized_full_prompt = self._tokenize_prompt(messages, self.hunyuanlarge_template, self.tokenizer.tokenizer)
461+ if self.args.append_eod:
462+ tokenized_full_prompt["input_ids"].append(self.tokenizer.eod)
463+ tokenized_full_prompt["attention_mask"].append(1)
464+ tokenized_full_prompt["labels"].append(self.tokenizer.eod)
465+ 
466+ for key in self.args.json_keys:
467+ tokenized_full_prompt[key] = [tokenized_full_prompt[key]]
468+ return tokenized_full_prompt
469+ 
470+ 
401class AlpacaStylePairwiseHandler(BaseDatasetHandler):471class AlpacaStylePairwiseHandler(BaseDatasetHandler):
402 """472 """
403 Handle alpaca style dataset format in pairwise dataset used in RM | DPO training473 Handle alpaca style dataset format in pairwise dataset used in RM | DPO training
@@ -983,6 +1053,7 @@ def build_dataset(args):
983 "AlpacaStylePairwiseHandler",1053 "AlpacaStylePairwiseHandler",
984 "SharegptStylePairwiseHandler",1054 "SharegptStylePairwiseHandler",
985 "PPOAlpacaStyleInstructionHandler",1055 "PPOAlpacaStyleInstructionHandler",
1056+ "HunyuanInstructionHandler",
986 "R1AlpacaStyleInstructionHandler",1057 "R1AlpacaStyleInstructionHandler",
987 "R1SharegptStyleInstructionHandler"1058 "R1SharegptStyleInstructionHandler"
988 ]:1059 ]:
@@ -194,7 +194,7 @@ class DecoderPackedMTFDataset(torch.utils.data.Dataset):
194 "labels": self._cut_token(item["labels"], np.int64),194 "labels": self._cut_token(item["labels"], np.int64),
195 "position_ids": self._cut_token(position_ids.numpy(), np.int64)195 "position_ids": self._cut_token(position_ids.numpy(), np.int64)
196 }196 }
197- elif self.args.stage == "prm":197+ elif self.args.stage == "prm" or self.args.cut_max_seqlen:
198 return {198 return {
199 "input_ids": self._cut_token(item['input_ids'], np.int64),199 "input_ids": self._cut_token(item['input_ids'], np.int64),
200 "attention_mask": self._cut_token(item["attention_mask"], np.int64),200 "attention_mask": self._cut_token(item["attention_mask"], np.int64),
@@ -76,7 +76,7 @@ def get_handler_dataset_attr(data_args, raw_datasets):
76 for column_name, target_name in data_args.map_keys.items():76 for column_name, target_name in data_args.map_keys.items():
77 setattr(dataset_attr, column_name, target_name)77 setattr(dataset_attr, column_name, target_name)
78 78 
79- elif "SharegptStyle" in data_args.handler_name:79+ elif "SharegptStyle" or "hunyuan" in data_args.handler_name:
80 dataset_attr.formatting = "sharegpt"80 dataset_attr.formatting = "sharegpt"
81 tag_names = ["role_tag", "content_tag", "user_tag", "assistant_tag", "observation_tag", "function_tag", "system_tag"]81 tag_names = ["role_tag", "content_tag", "user_tag", "assistant_tag", "observation_tag", "function_tag", "system_tag"]
82 column_names = ["messages", "tags", "system", "tools", "chosen", "rejected", "kto_tag"]82 column_names = ["messages", "tags", "system", "tools", "chosen", "rejected", "kto_tag"]
@@ -375,6 +375,95 @@ def convert_sharegpt_to_intermediate(
375 return outputs375 return outputs
376 376 
377 377 
378+def convert_hunyuan_to_intermediate(
379+ sample: Dict[str, List[Any]], dataset_attr: "InstructionDatasetAttr"):
380+ 
381+ outputs = {"prompt": [], "response": [], "system": [], "tools": []}
382+ 
383+ tag_mapping = {
384+ dataset_attr.user_tag: Role.USER.value,
385+ dataset_attr.assistant_tag: Role.ASSISTANT.value,
386+ dataset_attr.observation_tag: Role.OBSERVATION.value,
387+ dataset_attr.function_tag: Role.FUNCTION.value,
388+ dataset_attr.system_tag: Role.SYSTEM.value,
389+ }
390+ 
391+ # "human" and "observation" must appear in odd-numbered positions
392+ # "gpt" and "function" must appear in even-numbered positions.
393+ sys_tags = (dataset_attr.system_tag)
394+ odd_tags = (dataset_attr.user_tag, dataset_attr.observation_tag)
395+ even_tags = (dataset_attr.assistant_tag, dataset_attr.function_tag)
396+ accept_tags = (sys_tags, odd_tags, even_tags)
397+ 
398+ messages = sample[dataset_attr.messages]
399+ 
400+ if len(messages) == 0:
401+ return outputs
402+ 
403+ aligned_messages = []
404+ broken_data = False
405+ for turn_idx, message in enumerate(messages):
406+ if message[dataset_attr.role_tag] not in accept_tags[turn_idx % 3]:
407+ logger.warning("Invalid role tag in {}.".format(messages))
408+ broken_data = True
409+ 
410+ content_value = message.get(dataset_attr.content_tag)
411+ 
412+ if content_value is not None:
413+ aligned_messages.append(
414+ {"role": tag_mapping.get(message.get(dataset_attr.role_tag), "unknown"), "content": content_value}
415+ )
416+ else:
417+ logger.warning(f"Missing content tag in message at turn {turn_idx}: {message}")
418+ 
419+ is_message_count_divisible_by_3 = len(aligned_messages) % 3 == 0
420+ if (not dataset_attr.ranking and not is_message_count_divisible_by_3) or (
421+ dataset_attr.ranking and is_message_count_divisible_by_3
422+ ):
423+ logger.warning("Invalid message count in {}.".format(messages))
424+ broken_data = True
425+ 
426+ elif (
427+ dataset_attr.ranking
428+ and isinstance(sample[dataset_attr.chosen], dict)
429+ and isinstance(sample[dataset_attr.rejected], dict)
430+ ):
431+ chosen = sample[dataset_attr.chosen]
432+ rejected = sample[dataset_attr.rejected]
433+ if (
434+ chosen[dataset_attr.role_tag] not in accept_tags[-1]
435+ or rejected[dataset_attr.role_tag] not in accept_tags[-1]
436+ ):
437+ logger.warning("Invalid role tag in {}.".format([chosen, rejected]))
438+ broken_data = True
439+ 
440+ prompt = aligned_messages
441+ response = [
442+ {
443+ "role": tag_mapping.get(chosen.get(dataset_attr.role_tag, "gpt"), "assistant"),
444+ "content": chosen[dataset_attr.content_tag]
445+ },
446+ {
447+ "role": tag_mapping.get(rejected.get(dataset_attr.role_tag, "gpt"), "assistant"),
448+ "content": rejected[dataset_attr.content_tag]
449+ },
450+ ]
451+ 
452+ else: # normal example
453+ prompt = aligned_messages[:-1]
454+ response = aligned_messages[-1:]
455+ 
456+ if broken_data:
457+ logger.warning("Skipping this abnormal example.")
458+ return outputs
459+ 
460+ outputs["prompt"] = prompt
461+ outputs["response"] = response
462+ outputs["system"] = ""
463+ outputs["tools"].append(sample[dataset_attr.tools] if dataset_attr.tools else "")
464+ return outputs
465+ 
466+ 
378def align_dataset(dataset, dataset_attr, data_args):467def align_dataset(dataset, dataset_attr, data_args):
379 """468 """
380 Aligned dataset:469 Aligned dataset:
@@ -400,6 +489,8 @@ def align_dataset(dataset, dataset_attr, data_args):
400 """489 """
401 if dataset_attr.formatting == "alpaca":490 if dataset_attr.formatting == "alpaca":
402 convert_func = partial(convert_alpaca_to_intermediate, dataset_attr=dataset_attr)491 convert_func = partial(convert_alpaca_to_intermediate, dataset_attr=dataset_attr)
492+ elif data_args.handler_name == "HunyuanInstructionHandler":
493+ convert_func = partial(convert_hunyuan_to_intermediate, dataset_attr=dataset_attr)
403 else:494 else:
404 convert_func = partial(convert_sharegpt_to_intermediate, dataset_attr=dataset_attr)495 convert_func = partial(convert_sharegpt_to_intermediate, dataset_attr=dataset_attr)
405 496
@@ -58,10 +58,13 @@ def build_tokenizer(args):
58 tokenizer = TokenizerAdaptor(megatron_build_tokenizer(args))58 tokenizer = TokenizerAdaptor(megatron_build_tokenizer(args))
59 59 
60 if hasattr(args, "prompt_type") and args.prompt_type is not None:60 if hasattr(args, "prompt_type") and args.prompt_type is not None:
61- if ("PreTrainedTokenizerBase" not in str(tokenizer.tokenizer._pad.__func__)):61+ if hasattr(args, "handler_name") and args.handler_name == "HunyuanInstructionHandler":
62- tokenizer.tokenizer._pad = MethodType(PreTrainedTokenizerBase._pad, tokenizer.tokenizer)62+ pass
63- tokenizer.tokenizer.padding_side = "right"63+ else:
64- fix_model_tokenizer(tokenizer.tokenizer, args.prompt_type.strip(), args.prompt_type_path.strip())64+ if ("PreTrainedTokenizerBase" not in str(tokenizer.tokenizer._pad.__func__)):
65+ tokenizer.tokenizer._pad = MethodType(PreTrainedTokenizerBase._pad, tokenizer.tokenizer)
66+ tokenizer.tokenizer.padding_side = "right"
67+ fix_model_tokenizer(tokenizer.tokenizer, args.prompt_type.strip(), args.prompt_type_path.strip())
65 68 
66 return tokenizer69 return tokenizer
67 70 
@@ -113,7 +113,7 @@ def add_data_args(parser):
113 group.add_argument('--prompt-type', type=str, default=None,113 group.add_argument('--prompt-type', type=str, default=None,
114 choices=['default', 'empty', 'trl', 'chatglm2', 'chatglm3', 'chatglm3_system', 'glm4', 'chatml',114 choices=['default', 'empty', 'trl', 'chatglm2', 'chatglm3', 'chatglm3_system', 'glm4', 'chatml',
115 'chatml_de', 'qwen', 'qwen_r1', "qwen_math_r1", 'llama3', 'llama2', 'mistral', 'mixtral', 'gemma', 'alpaca',115 'chatml_de', 'qwen', 'qwen_r1', "qwen_math_r1", 'llama3', 'llama2', 'mistral', 'mixtral', 'gemma', 'alpaca',
116- 'deepseek2', 'deepseek2-lite', 'cpm', 'baichuan2', 'deepseek3', 'intern2'],116+ 'deepseek2', 'deepseek2-lite', 'cpm', 'baichuan2', 'deepseek3', 'intern2', 'hunyuan'],
117 help='Which template to use for constructing prompts in training.'117 help='Which template to use for constructing prompts in training.'
118 'e.g., "qwen"')118 'e.g., "qwen"')
119 group.add_argument('--prompt-type-path', type=str, default=TEMPLATES_DIR,119 group.add_argument('--prompt-type-path', type=str, default=TEMPLATES_DIR,
@@ -245,6 +245,7 @@ def validate_args(args):
245 "SharegptStylePairwiseHandler",245 "SharegptStylePairwiseHandler",
246 "AlpacaStyleProcessRewardHandler",246 "AlpacaStyleProcessRewardHandler",
247 "PPOAlpacaStyleInstructionHandler",247 "PPOAlpacaStyleInstructionHandler",
248+ "HunyuanInstructionHandler",
248 "R1AlpacaStyleInstructionHandler",249 "R1AlpacaStyleInstructionHandler",
249 "R1SharegptStyleInstructionHandler"250 "R1SharegptStyleInstructionHandler"
250 ]251 ]