已合并
【HunyuanLargeMoE】part of data-preprocess #2248
zhoubeirong创建于 2025年2月19日
【HunyuanLargeMoE】part of data-preprocess #2248
已合并
从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 | + | ||
| 401 | class AlpacaStylePairwiseHandler(BaseDatasetHandler): | 471 | class AlpacaStylePairwiseHandler(BaseDatasetHandler): |
| 402 | """ | 472 | """ |
| 403 | Handle alpaca style dataset format in pairwise dataset used in RM | DPO training | 473 | 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 outputs | 375 | 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 | + | ||
| 378 | def align_dataset(dataset, dataset_attr, data_args): | 467 | def 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 tokenizer | 69 | 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 | ] |