From 081bbee3666af10c3eb5654c27f150600c05c8d4 Mon Sep 17 00:00:00 2001 From: libo Date: Wed, 26 Feb 2025 10:50:40 +0000 Subject: [PATCH 1/4] init grpo --- .gitignore | 3 + .vscode/launch.json | 49 ++++++ .vscode/settings.json | 3 + dataset.py | 22 ++- ds_config_zero3.json | 3 +- requirements.txt | 7 +- reward_funcs.py | 20 +++ run.sh | 31 ++++ speech_grpo_trainer.py | 368 +++++++++++++++++++++++++++++++++++++++++ speech_llm.py | 19 ++- train.py | 32 +++- 11 files changed, 542 insertions(+), 15 deletions(-) create mode 100644 .gitignore create mode 100644 .vscode/launch.json create mode 100644 .vscode/settings.json create mode 100644 reward_funcs.py create mode 100755 run.sh create mode 100644 speech_grpo_trainer.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..74a7ff3 --- /dev/null +++ b/.gitignore @@ -0,0 +1,3 @@ +*pyc +*pth +*checkpoint* \ No newline at end of file diff --git a/.vscode/launch.json b/.vscode/launch.json new file mode 100644 index 0000000..b956e32 --- /dev/null +++ b/.vscode/launch.json @@ -0,0 +1,49 @@ +{ + "version": "0.2.0", + "configurations": [ + { + "name": "Train Qwen2-Whisper", + "type": "debugpy", + "module":"torch.distributed.launch", + "request": "launch", + "env": { + "HF_ENDPOINT": "https://hf-mirror.com", + "PYTHONPATH":"${workspaceRoot}:$PYTHONPATH", + "CUDA_VISIBLE_DEVICES":"1" + }, + "args": [ + "--use_env", + "--nnodes=1", + "--nproc_per_node=1", + "train.py", + "--llm_model_name_or_path", "Qwen/Qwen2.5-0.5B", + "--whisper_model_name_or_path", "tiny", + "--data_path", "/ceph2/user-data/linzhentao/code/ASR_training/wenet/examples/seewo/v3/data/edu_datas/zh_0403_in_wer/wer_0_5", + "--bf16", "True", + "--output_dir", "Qwen/Qwen2.5-0.5B-whisper-tiny", + "--num_train_epochs", "5", + "--per_device_train_batch_size", "4", + "--per_device_eval_batch_size", "1", + "--gradient_accumulation_steps", "2", + "--evaluation_strategy", "no", + "--save_strategy", "steps", + "--save_steps", "10000", + "--save_total_limit", "10", + "--learning_rate", "3e-4", + "--weight_decay", "0.01", + "--adam_beta2", "0.95", + "--warmup_ratio", "0.01", + "--lr_scheduler_type", "cosine", + "--logging_steps", "1", + "--report_to", "none", + "--model_max_length", "512", + "--gradient_checkpointing", + "--dataloader_num_workers", "4", + "--dataloader_prefetch_factor", "10", + "--deepspeed", "ds_config_zero3.json" + ], + "console": "integratedTerminal", + "justMyCode": false + } + ] +} diff --git a/.vscode/settings.json b/.vscode/settings.json new file mode 100644 index 0000000..3b66410 --- /dev/null +++ b/.vscode/settings.json @@ -0,0 +1,3 @@ +{ + "git.ignoreLimitWarning": true +} \ No newline at end of file diff --git a/dataset.py b/dataset.py index 02e8c11..25557df 100644 --- a/dataset.py +++ b/dataset.py @@ -12,6 +12,7 @@ import torchaudio import transformers import whisper +from tqdm import tqdm @dataclass @@ -40,9 +41,23 @@ def __init__( self.config = config self.inference = inference self.raw_data = [] + i = 0 with open(data_path, "r") as f: - for line in f: - self.raw_data.append(json.loads(line)) + for line in tqdm(f): + i += 1 + if i > 100000: + break + if not line.startswith('{'): + key, wav, txt, txt2, start, end, dur, u1, _, _, _ = line.split('\t') + obj = {} + obj['wav'] = wav + obj['key'] = key + obj['txt'] = txt.replace('▁',' ') + obj['start'] = round(float(start)) + obj['end'] = round(float(end)) + self.raw_data.append(obj) + else: + self.raw_data.append(json.loads(line)) def __len__(self): return len(self.raw_data) @@ -122,6 +137,9 @@ def __getitem__(self, i) -> Dict[str, torch.Tensor]: 'attention_mask': attention_mask, 'mel': mel, 'mel_len': mel_len, + 'prompt': 'Transcribe the speech', + 'key':msg['key'], + 'txt':msg['txt'] } if not self.inference: ret['labels'] = target_ids diff --git a/ds_config_zero3.json b/ds_config_zero3.json index e30fe94..1bfe6d6 100644 --- a/ds_config_zero3.json +++ b/ds_config_zero3.json @@ -25,7 +25,7 @@ "params": { "warmup_min_lr": "auto", "warmup_max_lr": "auto", - "warmup_num_steps": "auto" + "warmup_num_steps": 2500 } }, @@ -42,6 +42,7 @@ "overlap_comm": true, "contiguous_gradients": true, "sub_group_size": 1e9, + "reduce_scatter": true, "reduce_bucket_size": "auto", "stage3_prefetch_bucket_size": "auto", "stage3_param_persistence_threshold": "auto", diff --git a/requirements.txt b/requirements.txt index 9a81514..7ab4ae4 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,8 +1,9 @@ -deepspeed==0.14.4 +deepspeed==0.16.4 openai-whisper==20231117 peft==0.12.0 tensorboardX==2.6.2.2 torch>=2.2.2 torchaudio>=2.2.2 -transformers==4.43.3 -git+https://github.com/wenet-e2e/wenet.git +transformers==4.49.0 +trl>=0.16 +#git+https://github.com/wenet-e2e/wenet.git diff --git a/reward_funcs.py b/reward_funcs.py new file mode 100644 index 0000000..d8deff7 --- /dev/null +++ b/reward_funcs.py @@ -0,0 +1,20 @@ + +import jieba + +# for simple demo +# def reward_len(completions, **kwargs): +# return [-abs(20 - len(completion)) for completion in completions] + +def word_count(completions, **kwargs): + c = [] + for completion in completions: + words = jieba.cut(completion.strip()) + words = [w.strip() for w in words] + words = [w for w in words if w != ''] + print(words) + c.append(len(words)) + print('-'*120) + return c + # return [len(list(jieba.cut(completion))) for completion in completions] + +active_reward_func = [word_count] \ No newline at end of file diff --git a/run.sh b/run.sh new file mode 100755 index 0000000..1a58edd --- /dev/null +++ b/run.sh @@ -0,0 +1,31 @@ +export HF_ENDPOINT=https://hf-mirror.com + +DATA_PATH=/ceph2/user-data/linzhentao/code/ASR_training/wenet/examples/seewo/v3/data/edu_datas/zh_0403_in_wer/wer_0_5 +LLM_MODEL=Qwen/Qwen2-1.5B-Instruct +# LLM_MODEL=Qwen/Qwen2.5-0.5B +torchrun --standalone --nnodes=1 --nproc_per_node=1 train.py \ + --llm_model_name_or_path ${LLM_MODEL} \ + --whisper_model_name_or_path tiny \ + --data_path ${DATA_PATH} \ + --bf16 True \ + --output_dir ${LLM_MODEL}-whisper-tiny \ + --num_train_epochs 5 \ + --per_device_train_batch_size 8 \ + --per_device_eval_batch_size 1 \ + --gradient_accumulation_steps 8 \ + --evaluation_strategy "no" \ + --save_strategy "steps" \ + --save_steps 100 \ + --save_total_limit 10 \ + --learning_rate 3e-4 \ + --weight_decay 0.01 \ + --adam_beta2 0.95 \ + --warmup_ratio 0.01 \ + --lr_scheduler_type "cosine" \ + --logging_steps 1 \ + --report_to "none" \ + --model_max_length 512 \ + --gradient_checkpointing \ + --dataloader_num_workers 4 \ + --dataloader_prefetch_factor 10 \ + --deepspeed ds_config_zero3.json \ No newline at end of file diff --git a/speech_grpo_trainer.py b/speech_grpo_trainer.py new file mode 100644 index 0000000..14b65d4 --- /dev/null +++ b/speech_grpo_trainer.py @@ -0,0 +1,368 @@ +from typing import Any, Callable, Optional, Sized, Union +import torch +from torch import nn +from accelerate.utils import broadcast_object_list, gather, gather_object, is_peft_model, set_seed +from accelerate.utils.other import is_compiled_module +import transformers +from transformers import ( + AutoModelForCausalLM, + AutoModelForSequenceClassification, + AutoTokenizer, + GenerationConfig, + PreTrainedModel, + PreTrainedTokenizerBase, + Trainer, + TrainerCallback, + is_wandb_available, +) +from trl.data_utils import apply_chat_template, is_conversational, maybe_apply_chat_template +from trl import ModelConfig, GRPOConfig, GRPOTrainer +from trl.models import create_reference_model, prepare_deepspeed, unwrap_model_for_generation +from trl.import_utils import is_rich_available, is_vllm_available +from trl.trainer.utils import ( + generate_model_card, + get_comet_experiment_url, + pad, + print_prompt_completions_sample, + selective_log_softmax, +) +if is_wandb_available(): + import wandb + +class SpeechGRPOTrainer(GRPOTrainer): + def __init__( + self, + model: Union[str, PreTrainedModel], + reward_funcs, + args: Optional[GRPOConfig] = None, + train_dataset = None, + eval_dataset = None, + processing_class = None, + reward_processing_classes = None, + callbacks: Optional[list[TrainerCallback]] = None, + optimizers: tuple[Optional[torch.optim.Optimizer], Optional[torch.optim.lr_scheduler.LambdaLR]] = (None, None), + peft_config = None, + ): + super().__init__( + model=model, + reward_funcs = reward_funcs, + reward_processing_classes = reward_processing_classes, + args=args, + train_dataset=train_dataset, + eval_dataset=eval_dataset, + processing_class=processing_class, + callbacks=callbacks, + optimizers=optimizers, + peft_config = peft_config + ) + self.generation_config = GenerationConfig( + max_new_tokens=self.max_completion_length, + do_sample=True, + temperature=args.temperature, + pad_token_id=processing_class.pad_token_id, + num_beams=args.num_generations + ) + + + def _prepare_inputs(self, inputs: dict[str, Union[torch.Tensor, Any]]) -> dict[str, Union[torch.Tensor, Any]]: + mode = "eval" if self.control.should_evaluate else "train" + if mode == "train": + if self.state.global_step % self.num_iterations == 0: + inputs = self._generate_and_score_completions(inputs) + self._buffered_inputs[self._step % self.args.gradient_accumulation_steps] = inputs + else: + inputs = self._buffered_inputs[self._step % self.args.gradient_accumulation_steps] + self._step += 1 + else: + # In evaluation, we don't reuse completions across multiple updates, so we don't need to buffer inputs. + inputs = self._generate_and_score_completions(inputs) + return inputs + + def _generate_and_score_completions( + self, inputs: dict[str, Union[torch.Tensor, Any]] + ) -> dict[str, Union[torch.Tensor, Any]]: + device = self.accelerator.device + prompts = [x["prompt"] for x in inputs] + # prompts_text = [maybe_apply_chat_template(example, self.processing_class)["prompt"] for example in inputs] + # prompt_inputs = self.processing_class( + # prompts_text, return_tensors="pt", padding=True, padding_side="left", add_special_tokens=False + # ) + # prompt_inputs = super(GRPOTrainer, self)._prepare_inputs(prompt_inputs) + # prompt_ids, prompt_mask = prompt_inputs["input_ids"], prompt_inputs["attention_mask"] + + prompt_ids = torch.stack([x["input_ids"] for x in inputs]) + prompt_mask = torch.stack([x["attention_mask"] for x in inputs]) + mel = torch.stack([x["mel"] for x in inputs]) + mel_len = torch.stack([torch.tensor(x["mel_len"]) for x in inputs]) + + + if self.max_prompt_length is not None: + prompt_ids = prompt_ids[:, -self.max_prompt_length :] + prompt_mask = prompt_mask[:, -self.max_prompt_length :] + + # Generate completions using either vLLM or regular generation + if self.args.use_vllm: + # First, have main process load weights if needed + if self.state.global_step != self._last_loaded_step: + self._move_model_to_vllm() + self._last_loaded_step = self.state.global_step + + # Generate completions using vLLM: gather all prompts and use them in a single call in the main process + all_prompts_text = gather_object(prompts) + if self.accelerator.is_main_process: + # Since 'prompts' contains 'num_generations' duplicates, we first take unique prompts, and generate + # num_generations outputs for each one. This is faster than generating outputs for each duplicate + # prompt individually. + ordered_set_of_prompts = list(dict.fromkeys(all_prompts_text)) + all_outputs = self.llm.generate( + ordered_set_of_prompts, sampling_params=self.sampling_params, use_tqdm=False + ) + completion_ids = [] + for outputs in all_outputs: + for output in outputs.outputs: + completion_ids.append(output.token_ids) + else: + completion_ids = [None] * len(all_prompts_text) + # Broadcast the completions from the main process to all processes, ensuring each process receives its + # corresponding slice. + completion_ids = broadcast_object_list(completion_ids, from_process=0) + process_slice = slice( + self.accelerator.process_index * len(prompts), + (self.accelerator.process_index + 1) * len(prompts), + ) + completion_ids = completion_ids[process_slice] + + # Pad the completions, and concatenate them with the prompts + completion_ids = [torch.tensor(ids, device=device) for ids in completion_ids] + completion_ids = pad(completion_ids, padding_value=self.processing_class.pad_token_id) + prompt_completion_ids = torch.cat([prompt_ids, completion_ids], dim=1) + else: + # Regular generation path + # self.accelerator.free_memory() + # if decode_args.llm_type == 'qwen2': + eos_token_id = self.processing_class.convert_tokens_to_ids( + ['<|endoftext|>', '<|im_end|>']) + + with unwrap_model_for_generation(self.model, self.accelerator) as unwrapped_model: + prompt_completion_ids = unwrapped_model.generate( + prompt_ids, attention_mask=prompt_mask, + mel = mel, mel_len = mel_len, + eos_token_id=eos_token_id, + do_sample=True, + top_p=0.9, + temperature=0.7, + decode_config=self.generation_config, + ) + + # Compute prompt length and extract completion ids + # prompt_length = prompt_ids.size(1) + # prompt_ids = prompt_completion_ids[:, :prompt_length] + completion_ids = prompt_completion_ids#[:, prompt_length:] + + # Mask everything after the first EOS token + is_eos = completion_ids == self.processing_class.eos_token_id + eos_idx = torch.full((is_eos.size(0),), is_eos.size(1), dtype=torch.long, device=device) + eos_idx[is_eos.any(dim=1)] = is_eos.int().argmax(dim=1)[is_eos.any(dim=1)] + sequence_indices = torch.arange(is_eos.size(1), device=device).expand(is_eos.size(0), -1) + completion_mask = (sequence_indices <= eos_idx.unsqueeze(1)).int() + + # Concatenate prompt_mask with completion_mask for logit computation + attention_mask = torch.cat([prompt_mask, completion_mask], dim=1) # (B, P+C) + + logits_to_keep = completion_ids.size(1) # we only need to compute the logits for the completion tokens + + with torch.inference_mode(): + # When using num_iterations == 1, old_per_token_logps == per_token_logps, so we can skip it's + # computation here, and use per_token_logps.detach() instead. + if self.num_iterations > 1: + old_per_token_logps = self._get_per_token_logps( + self.model, prompt_completion_ids, attention_mask, mel, mel_len, logits_to_keep + ) + else: + old_per_token_logps = None + + if self.beta == 0.0: + ref_per_token_logps = None + elif self.ref_model is not None: + ref_per_token_logps = self._get_per_token_logps( + self.ref_model, prompt_completion_ids, attention_mask, mel, mel_len, logits_to_keep + ) + else: + with self.accelerator.unwrap_model(self.model).disable_adapter(): + ref_per_token_logps = self._get_per_token_logps( + self.model, prompt_completion_ids, attention_mask, mel, mel_len, logits_to_keep + ) + + # Decode the generated completions + completions_text = self.processing_class.batch_decode(completion_ids, skip_special_tokens=True) + # if is_conversational(inputs[0]): + # completions = [] + # for prompt, completion in zip(prompts, completions_text): + # bootstrap = prompt.pop()["content"] if prompt[-1]["role"] == "assistant" else "" + # completions.append([{"role": "assistant", "content": bootstrap + completion}]) + # else: + completions = completions_text + + rewards_per_func = torch.zeros(prompt_ids.size(0), len(self.reward_funcs), device=device) + for i, (reward_func, reward_processing_class) in enumerate( + zip(self.reward_funcs, self.reward_processing_classes) + ): + if isinstance(reward_func, nn.Module): # Module instead of PretrainedModel for compat with compiled models + # if is_conversational(inputs[0]): + # messages = [{"messages": p + c} for p, c in zip(prompts, completions)] + # texts = [apply_chat_template(x, reward_processing_class)["text"] for x in messages] + # else: + # texts = [p + c for p, c in zip(prompts, completions)] + texts = completions + reward_inputs = reward_processing_class( + texts, return_tensors="pt", padding=True, padding_side="right", add_special_tokens=False + ) + reward_inputs = super(GRPOTrainer, self)._prepare_inputs(reward_inputs) + with torch.inference_mode(): + rewards_per_func[:, i] = reward_func(**reward_inputs).logits[:, 0] # Shape (B*G,) + else: + # Repeat all input columns (but "prompt" and "completion") to match the number of generations + keys = [key for key in inputs[0] if key not in ["prompt", "completion"]] + reward_kwargs = {key: [example[key] for example in inputs] for key in keys} + output_reward_func = reward_func(prompts=prompts, completions=completions, **reward_kwargs) + rewards_per_func[:, i] = torch.tensor(output_reward_func, dtype=torch.float32, device=device) + + # Gather the reward per function: this part is crucial, because the rewards are normalized per group and the + # completions may be distributed across processes + rewards_per_func = gather(rewards_per_func) + + # Apply weights to each reward function's output and sum + rewards = (rewards_per_func * self.reward_weights.to(device).unsqueeze(0)).sum(dim=1) + + # Compute grouped-wise rewards + mean_grouped_rewards = rewards.view(-1, self.num_generations).mean(dim=1) + std_grouped_rewards = rewards.view(-1, self.num_generations).std(dim=1) + + # Normalize the rewards to compute the advantages + mean_grouped_rewards = mean_grouped_rewards.repeat_interleave(self.num_generations, dim=0) + std_grouped_rewards = std_grouped_rewards.repeat_interleave(self.num_generations, dim=0) + advantages = (rewards - mean_grouped_rewards) / (std_grouped_rewards + 1e-4) + + # Slice to keep only the local part of the data + process_slice = slice( + self.accelerator.process_index * len(prompts), + (self.accelerator.process_index + 1) * len(prompts), + ) + advantages = advantages[process_slice] + + # Log the metrics + mode = "eval" if self.control.should_evaluate else "train" + + completion_length = self.accelerator.gather_for_metrics(completion_mask.sum(1)).float().mean().item() + self._metrics[mode]["completion_length"].append(completion_length) + + reward_per_func = rewards_per_func.mean(0) + for i, reward_func in enumerate(self.reward_funcs): + if isinstance(reward_func, nn.Module): # Module instead of PretrainedModel for compat with compiled models + reward_func_name = reward_func.config._name_or_path.split("/")[-1] + else: + reward_func_name = reward_func.__name__ + self._metrics[mode][f"rewards/{reward_func_name}"].append(reward_per_func[i].item()) + + self._metrics[mode]["reward"].append(rewards.mean().item()) + self._metrics[mode]["reward_std"].append(std_grouped_rewards.mean().item()) + + if self.log_completions and self.state.global_step % self.args.logging_steps == 0: + # prompts_to_log = gather_object(prompts_text) + prompts_to_log = gather_object(prompts) + completions_to_log = gather_object(completions_text) + rewards_to_log = rewards.tolist() + + if self.accelerator.is_main_process: + if is_rich_available(): + print_prompt_completions_sample( + prompts_to_log, + completions_to_log, + rewards_to_log, + self.state.global_step, + ) + if self.args.report_to and "wandb" in self.args.report_to and wandb.run is not None: + import pandas as pd + + # For logging + table = { + "step": [str(self.state.global_step)] * len(rewards), + "prompt": prompts_to_log, + "completion": completions_to_log, + "reward": rewards.tolist(), + } + df = pd.DataFrame(table) + wandb.log({"completions": wandb.Table(dataframe=df)}) + + return { + "prompt_ids": prompt_ids, + "prompt_mask": prompt_mask, + "mel": mel, + "mel_len": mel_len, + "completion_ids": completion_ids, + "completion_mask": completion_mask, + "old_per_token_logps": old_per_token_logps, + "ref_per_token_logps": ref_per_token_logps, + "advantages": advantages, + } + def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None): + if return_outputs: + raise ValueError("The GRPOTrainer does not support returning outputs") + # Compute the per-token log probabilities for the model + + prompt_ids = inputs['prompt_ids'] # torch.stack([x["input_ids"] for x in inputs]) + prompt_mask = inputs['prompt_mask'] # torch.stack([x["attention_mask"] for x in inputs]) + mel = inputs['mel'] # torch.stack([x["mel"] for x in inputs]) + mel_len = inputs['mel_len'] # torch.stack([torch.tensor(x["mel_len"]) for x in inputs]) + + # prompt_ids, prompt_mask = inputs["prompt_ids"], inputs["prompt_mask"] + completion_ids, completion_mask = inputs["completion_ids"], inputs["completion_mask"] + input_ids = torch.cat([prompt_ids, completion_ids], dim=1) + attention_mask = torch.cat([prompt_mask, completion_mask], dim=1) + logits_to_keep = completion_ids.size(1) # we only need to compute the logits for the completion tokens + + per_token_logps = self._get_per_token_logps(model, input_ids, attention_mask, mel, mel_len, logits_to_keep) + + # Compute the KL divergence between the model and the reference model + if self.beta != 0.0: + ref_per_token_logps = inputs["ref_per_token_logps"] + per_token_kl = ( + torch.exp(ref_per_token_logps - per_token_logps) - (ref_per_token_logps - per_token_logps) - 1 + ) + + # Compute the loss + advantages = inputs["advantages"] + # When using num_iterations == 1, old_per_token_logps == per_token_logps, so we can skip it's computation (see + # _generate_and_score_completions) and use per_token_logps.detach() instead. + old_per_token_logps = inputs["old_per_token_logps"] if self.num_iterations > 1 else per_token_logps.detach() + coef_1 = torch.exp(per_token_logps - old_per_token_logps) + coef_2 = torch.clamp(coef_1, 1 - self.epsilon, 1 + self.epsilon) + per_token_loss1 = coef_1 * advantages.unsqueeze(1) + per_token_loss2 = coef_2 * advantages.unsqueeze(1) + per_token_loss = -torch.min(per_token_loss1, per_token_loss2) + if self.beta != 0.0: + per_token_loss = per_token_loss + self.beta * per_token_kl + loss = (per_token_loss * completion_mask).sum() / completion_mask.sum() + + # Log the metrics + mode = "eval" if self.control.should_evaluate else "train" + + if self.beta != 0.0: + mean_kl = ((per_token_kl * completion_mask).sum(dim=1) / completion_mask.sum(dim=1)).mean() + self._metrics[mode]["kl"].append(self.accelerator.gather_for_metrics(mean_kl).mean().item()) + + is_clipped = (per_token_loss1 < per_token_loss2).float() + clip_ratio = (is_clipped * completion_mask).sum() / completion_mask.sum() + self._metrics[mode]["clip_ratio"].append(self.accelerator.gather_for_metrics(clip_ratio).mean().item()) + return loss + + def _get_per_token_logps(self, model, input_ids, attention_mask, mel, mel_len, logits_to_keep): + # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded + logits = model(input_ids=input_ids, attention_mask=attention_mask, mel=mel, mel_len=mel_len, logits_to_keep=logits_to_keep + 1).logits + # logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred + + input_ids = input_ids[:, -logits_to_keep:] + # For transformers<=4.48, logits_to_keep argument isn't supported, so here we drop logits ourselves. + # See https://github.com/huggingface/trl/issues/2770 + logits = logits[:, -logits_to_keep:] + return selective_log_softmax(logits, input_ids) # compute logprobs for the input tokens diff --git a/speech_llm.py b/speech_llm.py index 1bb2401..32335f1 100644 --- a/speech_llm.py +++ b/speech_llm.py @@ -3,6 +3,7 @@ import math from typing import Optional from dataclasses import dataclass, field +from typing import Any, Callable, Optional, Sized, Union import safetensors import torch @@ -155,6 +156,7 @@ def forward( mel_len: torch.LongTensor = None, ctc_ids: torch.LongTensor = None, ctc_ids_len: torch.LongTensor = None, + logits_to_keep: Union[int, torch.Tensor] = 0 ): max_speech_size = self.model_args.max_speech_token_size text_emb = self.llm.get_input_embeddings()(input_ids) @@ -163,7 +165,8 @@ def forward( (speech_emb, text_emb[:, max_speech_size:, :]), dim=1) out = self.llm(inputs_embeds=inputs_embeds, attention_mask=attention_mask, - labels=labels) + labels=labels, + logits_to_keep=logits_to_keep) ctc_weight = self.model_args.ctc_weight if ctc_weight > 0: # Tie CTC linear transforme and input embedding weight @@ -172,7 +175,7 @@ def forward( ctc_act = ctc_act.transpose(0, 1) ctc_prob = ctc_act.log_softmax(2) prob_len = torch.ceil(mel_len / self.model_args.ds_rate).long() - with torch.cuda.amp.autocast(enabled=False): + with torch.amp.autocast(enabled=False): closs = self.ctc_loss(ctc_prob.float(), ctc_ids, prob_len, ctc_ids_len) out.loss = (1 - ctc_weight) * out.loss + ctc_weight * closs @@ -187,6 +190,10 @@ def generate( mel_len: torch.LongTensor = None, eos_token_id=None, decode_config=None, + do_sample=False, + top_p=1.0, + temperature=0.7 + **kwargs ): max_speech_size = self.model_args.max_speech_token_size text_emb = self.llm.get_input_embeddings()(input_ids) @@ -196,11 +203,15 @@ def generate( model_outputs = self.llm.generate( inputs_embeds=inputs_embeds, attention_mask=attention_mask, - do_sample=False, - top_p=1.0, + # do_sample=False, + # top_p=1.0, + do_sample=do_sample, + top_p=top_p, + temperature=temperature, num_beams=decode_config.num_beams, max_new_tokens=decode_config.max_new_tokens, eos_token_id=eos_token_id, + **kwargs ) return model_outputs diff --git a/train.py b/train.py index ed38d8c..37f9310 100644 --- a/train.py +++ b/train.py @@ -10,7 +10,9 @@ from dataset import DataArguments, SpeechDataset from speech_llm import init_model, ModelArguments - +from trl import GRPOConfig +from speech_grpo_trainer import SpeechGRPOTrainer +from reward_funcs import active_reward_func @dataclass class TrainingArguments(transformers.TrainingArguments): @@ -48,12 +50,32 @@ def main(): model_args) else: eval_dataset = None + + if training_args.grpo: + config = training_args.to_dict() + config['remove_unused_columns'] = False + config['num_generations'] = 4 #GRPO中的group number + grpo_config = GRPOConfig(**config) + trainer_cls = SpeechGRPOTrainer + trainer_arg = { + 'processing_class': tokenizer, + 'reward_funcs': active_reward_func + } + args = grpo_config + else: + trainer_cls = Trainer + trainer_arg = { + 'tokenizer': tokenizer, + } + args = training_args + + print(args.to_dict()) # Start trainer - trainer = Trainer(model=model, - tokenizer=tokenizer, - args=training_args, + trainer = trainer_cls(model=model, + args=grpo_config, train_dataset=train_dataset, - eval_dataset=eval_dataset) + eval_dataset=eval_dataset, + **trainer_arg) if list(pathlib.Path(training_args.output_dir).glob("checkpoint-*")): trainer.train(resume_from_checkpoint=True) else: From 366ebfbcf324d9ee11c503bb7062af40a38bc171 Mon Sep 17 00:00:00 2001 From: libo Date: Thu, 27 Feb 2025 03:11:14 +0000 Subject: [PATCH 2/4] fix some issue --- .gitignore | 4 +++- .vscode/launch.json | 8 +++++--- dataset.py | 11 ++++++++--- decode.sh | 7 +++++++ requirements.txt | 15 ++++++++++----- run.sh | 12 ++++++++---- speech_grpo_trainer.py | 13 ++++++++++--- speech_llm.py | 9 ++++----- train.py | 17 +++++++++++------ 9 files changed, 66 insertions(+), 30 deletions(-) create mode 100755 decode.sh diff --git a/.gitignore b/.gitignore index 74a7ff3..b7a80ec 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,5 @@ *pyc *pth -*checkpoint* \ No newline at end of file +*checkpoint* +*workspace +*txt \ No newline at end of file diff --git a/.vscode/launch.json b/.vscode/launch.json index b956e32..2abf044 100644 --- a/.vscode/launch.json +++ b/.vscode/launch.json @@ -16,11 +16,13 @@ "--nnodes=1", "--nproc_per_node=1", "train.py", - "--llm_model_name_or_path", "Qwen/Qwen2.5-0.5B", + "--grpo", + "--llm_model_name_or_path", "Qwen/Qwen2-1.5B-Instruct", "--whisper_model_name_or_path", "tiny", - "--data_path", "/ceph2/user-data/linzhentao/code/ASR_training/wenet/examples/seewo/v3/data/edu_datas/zh_0403_in_wer/wer_0_5", + "--data_path", "/ceph2/user-data/chenzhongliang/west/aishell/train.jsonl", "--bf16", "True", - "--output_dir", "Qwen/Qwen2.5-0.5B-whisper-tiny", + "--projector_model_path", "/ceph2/user-data/chenzhongliang/west/Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors", + "--output_dir", "Qwen-1.5B-Instruct-whisper-tiny", "--num_train_epochs", "5", "--per_device_train_batch_size", "4", "--per_device_eval_batch_size", "1", diff --git a/dataset.py b/dataset.py index 25557df..1441110 100644 --- a/dataset.py +++ b/dataset.py @@ -34,6 +34,7 @@ def __init__( tokenizer: transformers.PreTrainedTokenizer, config, # model config inference: bool = False, + grpo = False, ): super(SpeechDataset, self).__init__() print("Formatting inputs...") @@ -58,6 +59,7 @@ def __init__( self.raw_data.append(obj) else: self.raw_data.append(json.loads(line)) + self.grpo = grpo def __len__(self): return len(self.raw_data) @@ -137,10 +139,13 @@ def __getitem__(self, i) -> Dict[str, torch.Tensor]: 'attention_mask': attention_mask, 'mel': mel, 'mel_len': mel_len, - 'prompt': 'Transcribe the speech', - 'key':msg['key'], - 'txt':msg['txt'] } + if self.grpo: + ret['prompt'] = instruction + if 'key' in msg: + ret['key'] = msg['key'] + ret['txt'] = msg['txt'] + if not self.inference: ret['labels'] = target_ids ret['ctc_ids'] = ctc_ids diff --git a/decode.sh b/decode.sh new file mode 100755 index 0000000..bd15dcb --- /dev/null +++ b/decode.sh @@ -0,0 +1,7 @@ +python recognize.py \ + --llm_model_name_or_path Qwen/Qwen2-1.5B-Instruct \ + --whisper_model_name_or_path tiny \ + --projector_model_path /ceph2/user-data/chenzhongliang/west/Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors \ + --data_path /ceph2/user-data/chenzhongliang/west/aishell/test.jsonl \ + --result_path result.txt +` \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 7ab4ae4..2331e43 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,9 +1,14 @@ deepspeed==0.16.4 -openai-whisper==20231117 -peft==0.12.0 +openai-whisper +#==20231117 tensorboardX==2.6.2.2 -torch>=2.2.2 -torchaudio>=2.2.2 +#torch +#>=2.2.2 +#peft>=0.12.0 +#torchaudio +#>=2.2.2 transformers==4.49.0 -trl>=0.16 +accelerate>=1.4.0 +#trl>=0.16 #git+https://github.com/wenet-e2e/wenet.git +jieba \ No newline at end of file diff --git a/run.sh b/run.sh index 1a58edd..43ac64d 100755 --- a/run.sh +++ b/run.sh @@ -1,16 +1,20 @@ export HF_ENDPOINT=https://hf-mirror.com DATA_PATH=/ceph2/user-data/linzhentao/code/ASR_training/wenet/examples/seewo/v3/data/edu_datas/zh_0403_in_wer/wer_0_5 +DATA_PATH=/ceph2/user-data/chenzhongliang/west/aishell/train.jsonl LLM_MODEL=Qwen/Qwen2-1.5B-Instruct # LLM_MODEL=Qwen/Qwen2.5-0.5B -torchrun --standalone --nnodes=1 --nproc_per_node=1 train.py \ +PER_DEVICE_TRAIN_BATCH_SIZE=96 +torchrun --standalone --nnodes=1 --nproc_per_node=8 train.py \ + --grpo \ + --projector_model_path Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors \ --llm_model_name_or_path ${LLM_MODEL} \ --whisper_model_name_or_path tiny \ --data_path ${DATA_PATH} \ --bf16 True \ --output_dir ${LLM_MODEL}-whisper-tiny \ --num_train_epochs 5 \ - --per_device_train_batch_size 8 \ + --per_device_train_batch_size ${PER_DEVICE_TRAIN_BATCH_SIZE} \ --per_device_eval_batch_size 1 \ --gradient_accumulation_steps 8 \ --evaluation_strategy "no" \ @@ -26,6 +30,6 @@ torchrun --standalone --nnodes=1 --nproc_per_node=1 train.py \ --report_to "none" \ --model_max_length 512 \ --gradient_checkpointing \ - --dataloader_num_workers 4 \ - --dataloader_prefetch_factor 10 \ + --dataloader_num_workers 16 \ + --dataloader_prefetch_factor 64 \ --deepspeed ds_config_zero3.json \ No newline at end of file diff --git a/speech_grpo_trainer.py b/speech_grpo_trainer.py index 14b65d4..1a71506 100644 --- a/speech_grpo_trainer.py +++ b/speech_grpo_trainer.py @@ -149,8 +149,11 @@ def _generate_and_score_completions( mel = mel, mel_len = mel_len, eos_token_id=eos_token_id, do_sample=True, - top_p=0.9, - temperature=0.7, + repetition_penalty=1.2, + no_repeat_ngram_size=3, + + # top_p=0.9, + temperature=0.9, decode_config=self.generation_config, ) @@ -358,7 +361,11 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N def _get_per_token_logps(self, model, input_ids, attention_mask, mel, mel_len, logits_to_keep): # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded - logits = model(input_ids=input_ids, attention_mask=attention_mask, mel=mel, mel_len=mel_len, logits_to_keep=logits_to_keep + 1).logits + logits = model(input_ids=input_ids, + attention_mask=attention_mask, + mel=mel, + mel_len=mel_len, + logits_to_keep=logits_to_keep + 1).logits # logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred input_ids = input_ids[:, -logits_to_keep:] diff --git a/speech_llm.py b/speech_llm.py index 32335f1..9cd87ee 100644 --- a/speech_llm.py +++ b/speech_llm.py @@ -119,8 +119,8 @@ def __init__( # Do not save the parameter of llm and whisper for k in self.llm.state_dict().keys(): self._keys_to_ignore_on_save.add('llm.' + k) - for k in self.encoder.state_dict().keys(): - self._keys_to_ignore_on_save.add('encoder.' + k) + # for k in self.encoder.state_dict().keys(): + # self._keys_to_ignore_on_save.add('encoder.' + k) # Use bos_token_id as CTC blank id self.ctc_loss = nn.CTCLoss(config.bos_token_id, reduction='mean', @@ -192,7 +192,7 @@ def generate( decode_config=None, do_sample=False, top_p=1.0, - temperature=0.7 + temperature=0.7, **kwargs ): max_speech_size = self.model_args.max_speech_token_size @@ -203,14 +203,13 @@ def generate( model_outputs = self.llm.generate( inputs_embeds=inputs_embeds, attention_mask=attention_mask, - # do_sample=False, - # top_p=1.0, do_sample=do_sample, top_p=top_p, temperature=temperature, num_beams=decode_config.num_beams, max_new_tokens=decode_config.max_new_tokens, eos_token_id=eos_token_id, + pad_token_id=decode_config.pad_token_id, **kwargs ) return model_outputs diff --git a/train.py b/train.py index 37f9310..8febb25 100644 --- a/train.py +++ b/train.py @@ -17,6 +17,8 @@ @dataclass class TrainingArguments(transformers.TrainingArguments): optim: str = field(default="adafactor") + grpo: bool = field(default=False) + remove_unused_columns: bool = field(default=False) def main(): @@ -44,7 +46,7 @@ def main(): tokenizer.pad_token = '<|finetune_right_pad_id|>' print("Loading data...") - train_dataset = SpeechDataset(data_args.data_path, tokenizer, model_args) + train_dataset = SpeechDataset(data_args.data_path, tokenizer, model_args, grpo = training_args.grpo) if data_args.eval_data_path: eval_dataset = SpeechDataset(data_args.eval_data_path, tokenizer, model_args) @@ -54,7 +56,8 @@ def main(): if training_args.grpo: config = training_args.to_dict() config['remove_unused_columns'] = False - config['num_generations'] = 4 #GRPO中的group number + config['num_generations'] = 2 #GRPO中的group number + del config['grpo'] grpo_config = GRPOConfig(**config) trainer_cls = SpeechGRPOTrainer trainer_arg = { @@ -67,18 +70,20 @@ def main(): trainer_arg = { 'tokenizer': tokenizer, } + # training_args.remove_unused_columns = False args = training_args print(args.to_dict()) # Start trainer trainer = trainer_cls(model=model, - args=grpo_config, + args=args, train_dataset=train_dataset, eval_dataset=eval_dataset, **trainer_arg) - if list(pathlib.Path(training_args.output_dir).glob("checkpoint-*")): - trainer.train(resume_from_checkpoint=True) - else: + if 1: + # if list(pathlib.Path(training_args.output_dir).glob("checkpoint-*")): + # trainer.train(resume_from_checkpoint=True) + # else: trainer.train() trainer.save_state() From efbc8ef3a861d23ac98e418c22086d17ee204d2c Mon Sep 17 00:00:00 2001 From: libo Date: Fri, 28 Feb 2025 09:17:35 +0000 Subject: [PATCH 3/4] fix ref_model source, from pretrained speechLM --- .gitignore | 7 +- .vscode/launch.json | 46 +++++++++++- dataset.py | 3 + decode.sh | 7 -- ds_config_zero3.json | 6 +- export.py | 24 ++++++ recognize.py | 17 ++++- requirements.txt | 3 +- reward_funcs.py | 22 ++++-- run.sh | 109 ++++++++++++++++++++------- speech_grpo_trainer.py | 108 ++++++++------------------- speech_llm.py | 166 +++++++++++++++++++++++++++++------------ train.py | 13 +++- 13 files changed, 353 insertions(+), 178 deletions(-) delete mode 100755 decode.sh create mode 100644 export.py diff --git a/.gitignore b/.gitignore index b7a80ec..c757b5a 100644 --- a/.gitignore +++ b/.gitignore @@ -2,4 +2,9 @@ *pth *checkpoint* *workspace -*txt \ No newline at end of file +*txt +core* +*safetensors +*swp +*jsonl +west-slm/ \ No newline at end of file diff --git a/.vscode/launch.json b/.vscode/launch.json index 2abf044..d6c578a 100644 --- a/.vscode/launch.json +++ b/.vscode/launch.json @@ -19,10 +19,11 @@ "--grpo", "--llm_model_name_or_path", "Qwen/Qwen2-1.5B-Instruct", "--whisper_model_name_or_path", "tiny", + "--temperature","0.5", "--data_path", "/ceph2/user-data/chenzhongliang/west/aishell/train.jsonl", "--bf16", "True", - "--projector_model_path", "/ceph2/user-data/chenzhongliang/west/Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors", - "--output_dir", "Qwen-1.5B-Instruct-whisper-tiny", + "--projector_model_path", "Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors", + "--output_dir", "Qwen/Qwen2-1.5B-Instruct-whisper-tiny", "--num_train_epochs", "5", "--per_device_train_batch_size", "4", "--per_device_eval_batch_size", "1", @@ -46,6 +47,47 @@ ], "console": "integratedTerminal", "justMyCode": false + }, + { + "name": "export", + "type": "debugpy", + "module":"torch.distributed.launch", + "request": "launch", + "env": { + "HF_ENDPOINT": "https://hf-mirror.com", + "PYTHONPATH":"${workspaceRoot}:$PYTHONPATH", + "CUDA_VISIBLE_DEVICES":"1" + }, + "args": [ + "export.py", + "--llm_model_name_or_path", "Qwen/Qwen2-1.5B-Instruct", + "--whisper_model_name_or_path", "tiny", + "--data_path", "/ceph2/user-data/chenzhongliang/west/aishell/train.jsonl", + "--bf16", "True", + "--projector_model_path", "Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors", + "--output_dir", "Qwen/Qwen2-1.5B-Instruct-whisper-tiny", + "--num_train_epochs", "5", + "--per_device_train_batch_size", "4", + "--per_device_eval_batch_size", "1", + "--gradient_accumulation_steps", "2", + "--evaluation_strategy", "no", + "--save_strategy", "steps", + "--save_steps", "10000", + "--save_total_limit", "10", + "--learning_rate", "3e-4", + "--weight_decay", "0.01", + "--adam_beta2", "0.95", + "--warmup_ratio", "0.01", + "--lr_scheduler_type", "cosine", + "--logging_steps", "1", + "--report_to", "none", + "--model_max_length", "512", + "--gradient_checkpointing", + "--dataloader_num_workers", "4", + "--dataloader_prefetch_factor", "10", + ], + "console": "integratedTerminal", + "justMyCode": false } ] } diff --git a/dataset.py b/dataset.py index 1441110..01fdcb0 100644 --- a/dataset.py +++ b/dataset.py @@ -60,6 +60,8 @@ def __init__( else: self.raw_data.append(json.loads(line)) self.grpo = grpo + if self.grpo: + self.inference = True def __len__(self): return len(self.raw_data) @@ -145,6 +147,7 @@ def __getitem__(self, i) -> Dict[str, torch.Tensor]: if 'key' in msg: ret['key'] = msg['key'] ret['txt'] = msg['txt'] + ret['wav'] = msg['wav'] if not self.inference: ret['labels'] = target_ids diff --git a/decode.sh b/decode.sh deleted file mode 100755 index bd15dcb..0000000 --- a/decode.sh +++ /dev/null @@ -1,7 +0,0 @@ -python recognize.py \ - --llm_model_name_or_path Qwen/Qwen2-1.5B-Instruct \ - --whisper_model_name_or_path tiny \ - --projector_model_path /ceph2/user-data/chenzhongliang/west/Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors \ - --data_path /ceph2/user-data/chenzhongliang/west/aishell/test.jsonl \ - --result_path result.txt -` \ No newline at end of file diff --git a/ds_config_zero3.json b/ds_config_zero3.json index 1bfe6d6..4a9e5b3 100644 --- a/ds_config_zero3.json +++ b/ds_config_zero3.json @@ -39,13 +39,13 @@ "device": "none", "pin_memory": true }, - "overlap_comm": true, + "overlap_comm": false, "contiguous_gradients": true, "sub_group_size": 1e9, - "reduce_scatter": true, + "reduce_scatter": false, "reduce_bucket_size": "auto", "stage3_prefetch_bucket_size": "auto", - "stage3_param_persistence_threshold": "auto", + "stage3_param_persistence_threshold": 1e10, "stage3_max_live_parameters": 1e9, "stage3_max_reuse_distance": 1e9, "stage3_gather_16bit_weights_on_model_save": true diff --git a/export.py b/export.py new file mode 100644 index 0000000..7f9621c --- /dev/null +++ b/export.py @@ -0,0 +1,24 @@ +from accelerate import Accelerator +from speech_llm import init_model, ModelArguments +import transformers +from dataset import DataArguments +from train import TrainingArguments +from trl.models import unwrap_model_for_generation + +def export_model(model, output_dir): + accelerator = Accelerator() + with unwrap_model_for_generation(model, accelerator) as unwrapped_model: + unwrapped_model.save_pretrained(output_dir) + +if __name__ == '__main__': + parser = transformers.HfArgumentParser( + (ModelArguments, DataArguments,TrainingArguments )) + ( + model_args, + data_args, + _ + ) = parser.parse_args_into_dataclasses() + + model = init_model(model_args) + model.freeze_llm() + export_model(model, './west-slm') \ No newline at end of file diff --git a/recognize.py b/recognize.py index ecdb5f3..e3cff4b 100644 --- a/recognize.py +++ b/recognize.py @@ -11,7 +11,8 @@ from dataset import SpeechDataset, DataArguments from speech_llm import init_model, ModelArguments - +from transformers import GenerationConfig + @dataclass class DecodeArguments: @@ -32,6 +33,7 @@ def main(): if decode_args.llm_type == 'qwen2': eos_token_id = tokenizer.convert_tokens_to_ids( ['<|endoftext|>', '<|im_end|>']) + decode_args.pad_token_id = tokenizer.pad_token_id else: tokenizer.pad_token = '<|finetune_right_pad_id|>' eos_token_id = tokenizer.convert_tokens_to_ids( @@ -52,11 +54,20 @@ def main(): decode_func = model.generate else: decode_func = model.decode_ctc + generation_config = GenerationConfig( + do_sample=False, + pad_token_id=tokenizer.pad_token_id, + eos_token_id=tokenizer.eos_token_id, + max_new_tokens=100, + num_beams=1, + ) with torch.no_grad(): for item in tqdm(data_loader): generated_ids = decode_func(**item, - eos_token_id=eos_token_id, - decode_config=decode_args) + decode_config=generation_config, + repetition_penalty=1.2, + no_repeat_ngram_size=3, + ) text = tokenizer.batch_decode(generated_ids, skip_special_tokens=True) print(text) diff --git a/requirements.txt b/requirements.txt index 2331e43..1f5f3b9 100644 --- a/requirements.txt +++ b/requirements.txt @@ -11,4 +11,5 @@ transformers==4.49.0 accelerate>=1.4.0 #trl>=0.16 #git+https://github.com/wenet-e2e/wenet.git -jieba \ No newline at end of file +jieba +editdistance \ No newline at end of file diff --git a/reward_funcs.py b/reward_funcs.py index d8deff7..a9f270e 100644 --- a/reward_funcs.py +++ b/reward_funcs.py @@ -1,20 +1,30 @@ import jieba - +import editdistance +import numpy as np # for simple demo # def reward_len(completions, **kwargs): # return [-abs(20 - len(completion)) for completion in completions] +def editdistance_score(completions, **kwargs): + # return [0] * len(completions) + diff = [] + for hyp, lab in zip(completions,kwargs['txt']): + # diff.append(-1.0 * editdistance.eval(hyp, lab) / len(lab)) + diff.append(np.log(1e-9 + editdistance.eval(hyp, lab))) + return diff + def word_count(completions, **kwargs): c = [] - for completion in completions: + for i, completion in enumerate(completions): words = jieba.cut(completion.strip()) words = [w.strip() for w in words] words = [w for w in words if w != ''] - print(words) - c.append(len(words)) - print('-'*120) + if i==0: + print(kwargs['wav'][i],words) + c.append(-abs(0 - len(completion))) + # print('-'*120) return c # return [len(list(jieba.cut(completion))) for completion in completions] -active_reward_func = [word_count] \ No newline at end of file +active_reward_func = [editdistance_score] \ No newline at end of file diff --git a/run.sh b/run.sh index 43ac64d..eaf818e 100755 --- a/run.sh +++ b/run.sh @@ -4,32 +4,83 @@ DATA_PATH=/ceph2/user-data/linzhentao/code/ASR_training/wenet/examples/seewo/v3/ DATA_PATH=/ceph2/user-data/chenzhongliang/west/aishell/train.jsonl LLM_MODEL=Qwen/Qwen2-1.5B-Instruct # LLM_MODEL=Qwen/Qwen2.5-0.5B -PER_DEVICE_TRAIN_BATCH_SIZE=96 -torchrun --standalone --nnodes=1 --nproc_per_node=8 train.py \ - --grpo \ - --projector_model_path Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors \ - --llm_model_name_or_path ${LLM_MODEL} \ - --whisper_model_name_or_path tiny \ - --data_path ${DATA_PATH} \ - --bf16 True \ - --output_dir ${LLM_MODEL}-whisper-tiny \ - --num_train_epochs 5 \ - --per_device_train_batch_size ${PER_DEVICE_TRAIN_BATCH_SIZE} \ - --per_device_eval_batch_size 1 \ - --gradient_accumulation_steps 8 \ - --evaluation_strategy "no" \ - --save_strategy "steps" \ - --save_steps 100 \ - --save_total_limit 10 \ - --learning_rate 3e-4 \ - --weight_decay 0.01 \ - --adam_beta2 0.95 \ - --warmup_ratio 0.01 \ - --lr_scheduler_type "cosine" \ - --logging_steps 1 \ - --report_to "none" \ - --model_max_length 512 \ - --gradient_checkpointing \ - --dataloader_num_workers 16 \ - --dataloader_prefetch_factor 64 \ - --deepspeed ds_config_zero3.json \ No newline at end of file +PER_DEVICE_TRAIN_BATCH_SIZE=48 +stage=$1 +stop_stage=$1 + +if [ ${stage} -le 0 ] && [ ${stop_stage} -ge 0 ]; then + torchrun --standalone --nnodes=1 --nproc_per_node=8 train.py \ + --projector_model_path Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors \ + --llm_model_name_or_path ${LLM_MODEL} \ + --whisper_model_name_or_path tiny \ + --data_path ${DATA_PATH} \ + --bf16 True \ + --output_dir ${LLM_MODEL}-whisper-tiny \ + --num_train_epochs 5 \ + --per_device_train_batch_size ${PER_DEVICE_TRAIN_BATCH_SIZE} \ + --per_device_eval_batch_size 1 \ + --gradient_accumulation_steps 8 \ + --evaluation_strategy "no" \ + --save_strategy "steps" \ + --save_steps 100 \ + --save_total_limit 10 \ + --learning_rate 3e-4 \ + --weight_decay 0.01 \ + --adam_beta2 0.95 \ + --warmup_ratio 0.01 \ + --lr_scheduler_type "cosine" \ + --logging_steps 1 \ + --report_to "none" \ + --model_max_length 512 \ + --gradient_checkpointing \ + --dataloader_num_workers 16 \ + --dataloader_prefetch_factor 64 \ + --deepspeed ds_config_zero3.json + +fi + + +if [ ${stage} -le 1 ] && [ ${stop_stage} -ge 1 ]; then + torchrun --standalone --nnodes=1 --nproc_per_node=8 train.py \ + --grpo \ + --projector_model_path Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors \ + --llm_model_name_or_path ${LLM_MODEL} \ + --whisper_model_name_or_path tiny \ + --data_path ${DATA_PATH} \ + --bf16 True \ + --output_dir ${LLM_MODEL}-whisper-tiny \ + --temperature 0.5 \ + --num_train_epochs 5 \ + --per_device_train_batch_size ${PER_DEVICE_TRAIN_BATCH_SIZE} \ + --per_device_eval_batch_size 1 \ + --gradient_accumulation_steps 8 \ + --evaluation_strategy "no" \ + --save_strategy "steps" \ + --save_steps 1 \ + --save_total_limit 10 \ + --learning_rate 3e-4 \ + --weight_decay 0.01 \ + --adam_beta2 0.95 \ + --warmup_ratio 0.01 \ + --lr_scheduler_type "cosine" \ + --logging_steps 1 \ + --report_to "none" \ + --model_max_length 512 \ + --gradient_checkpointing \ + --dataloader_num_workers 16 \ + --dataloader_prefetch_factor 64 \ + --deepspeed ds_config_zero3.json + + +fi + +TEST_MODEL=Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors +TEST_MODEL=Qwen/Qwen2-1.5B-Instruct-whisper-tiny/checkpoint-3/model.safetensors +if [ ${stage} -le 2 ] && [ ${stop_stage} -ge 2 ]; then + python recognize.py \ + --llm_model_name_or_path ${LLM_MODEL} \ + --whisper_model_name_or_path tiny \ + --projector_model_path ${TEST_MODEL} \ + --data_path test.jsonl \ + --result_path result.txt +fi \ No newline at end of file diff --git a/speech_grpo_trainer.py b/speech_grpo_trainer.py index 1a71506..fb9e291 100644 --- a/speech_grpo_trainer.py +++ b/speech_grpo_trainer.py @@ -1,31 +1,24 @@ from typing import Any, Callable, Optional, Sized, Union import torch from torch import nn -from accelerate.utils import broadcast_object_list, gather, gather_object, is_peft_model, set_seed -from accelerate.utils.other import is_compiled_module +from accelerate.utils import broadcast_object_list, gather, gather_object import transformers from transformers import ( - AutoModelForCausalLM, - AutoModelForSequenceClassification, - AutoTokenizer, GenerationConfig, PreTrainedModel, - PreTrainedTokenizerBase, - Trainer, TrainerCallback, is_wandb_available, ) -from trl.data_utils import apply_chat_template, is_conversational, maybe_apply_chat_template -from trl import ModelConfig, GRPOConfig, GRPOTrainer -from trl.models import create_reference_model, prepare_deepspeed, unwrap_model_for_generation -from trl.import_utils import is_rich_available, is_vllm_available +from trl import GRPOConfig, GRPOTrainer +from trl.models import unwrap_model_for_generation +from trl.import_utils import is_rich_available +from trl.extras.profiling import profiling_decorator from trl.trainer.utils import ( - generate_model_card, - get_comet_experiment_url, pad, print_prompt_completions_sample, selective_log_softmax, ) +from copy import deepcopy if is_wandb_available(): import wandb @@ -42,6 +35,7 @@ def __init__( callbacks: Optional[list[TrainerCallback]] = None, optimizers: tuple[Optional[torch.optim.Optimizer], Optional[torch.optim.lr_scheduler.LambdaLR]] = (None, None), peft_config = None, + ref_model = None, ): super().__init__( model=model, @@ -56,40 +50,25 @@ def __init__( peft_config = peft_config ) self.generation_config = GenerationConfig( - max_new_tokens=self.max_completion_length, do_sample=True, temperature=args.temperature, pad_token_id=processing_class.pad_token_id, - num_beams=args.num_generations + eos_token_id=processing_class.eos_token_id, + max_new_tokens=100, + num_beams=1, + top_p=0.9, ) + if ref_model is not None: + self.ref_model = ref_model + # with unwrap_model_for_generation(ref_model, self.accelerator) as unwrapped_model: + # self.ref_model = deepcopy(unwrapped_model) - - def _prepare_inputs(self, inputs: dict[str, Union[torch.Tensor, Any]]) -> dict[str, Union[torch.Tensor, Any]]: - mode = "eval" if self.control.should_evaluate else "train" - if mode == "train": - if self.state.global_step % self.num_iterations == 0: - inputs = self._generate_and_score_completions(inputs) - self._buffered_inputs[self._step % self.args.gradient_accumulation_steps] = inputs - else: - inputs = self._buffered_inputs[self._step % self.args.gradient_accumulation_steps] - self._step += 1 - else: - # In evaluation, we don't reuse completions across multiple updates, so we don't need to buffer inputs. - inputs = self._generate_and_score_completions(inputs) - return inputs - + @profiling_decorator def _generate_and_score_completions( self, inputs: dict[str, Union[torch.Tensor, Any]] ) -> dict[str, Union[torch.Tensor, Any]]: device = self.accelerator.device prompts = [x["prompt"] for x in inputs] - # prompts_text = [maybe_apply_chat_template(example, self.processing_class)["prompt"] for example in inputs] - # prompt_inputs = self.processing_class( - # prompts_text, return_tensors="pt", padding=True, padding_side="left", add_special_tokens=False - # ) - # prompt_inputs = super(GRPOTrainer, self)._prepare_inputs(prompt_inputs) - # prompt_ids, prompt_mask = prompt_inputs["input_ids"], prompt_inputs["attention_mask"] - prompt_ids = torch.stack([x["input_ids"] for x in inputs]) prompt_mask = torch.stack([x["attention_mask"] for x in inputs]) mel = torch.stack([x["mel"] for x in inputs]) @@ -138,29 +117,17 @@ def _generate_and_score_completions( prompt_completion_ids = torch.cat([prompt_ids, completion_ids], dim=1) else: # Regular generation path - # self.accelerator.free_memory() - # if decode_args.llm_type == 'qwen2': - eos_token_id = self.processing_class.convert_tokens_to_ids( - ['<|endoftext|>', '<|im_end|>']) - + with unwrap_model_for_generation(self.model, self.accelerator) as unwrapped_model: - prompt_completion_ids = unwrapped_model.generate( + completion_ids = unwrapped_model.generate( prompt_ids, attention_mask=prompt_mask, mel = mel, mel_len = mel_len, - eos_token_id=eos_token_id, - do_sample=True, repetition_penalty=1.2, no_repeat_ngram_size=3, - - # top_p=0.9, - temperature=0.9, decode_config=self.generation_config, ) - # Compute prompt length and extract completion ids - # prompt_length = prompt_ids.size(1) - # prompt_ids = prompt_completion_ids[:, :prompt_length] - completion_ids = prompt_completion_ids#[:, prompt_length:] + prompt_completion_ids = torch.cat([prompt_ids, completion_ids], dim=1) # Mask everything after the first EOS token is_eos = completion_ids == self.processing_class.eos_token_id @@ -187,9 +154,10 @@ def _generate_and_score_completions( if self.beta == 0.0: ref_per_token_logps = None elif self.ref_model is not None: - ref_per_token_logps = self._get_per_token_logps( - self.ref_model, prompt_completion_ids, attention_mask, mel, mel_len, logits_to_keep - ) + with unwrap_model_for_generation(self.ref_model, self.accelerator) as unwrapped_model: + ref_per_token_logps = self._get_per_token_logps( + unwrapped_model, prompt_completion_ids, attention_mask, mel, mel_len, logits_to_keep + ) else: with self.accelerator.unwrap_model(self.model).disable_adapter(): ref_per_token_logps = self._get_per_token_logps( @@ -198,25 +166,13 @@ def _generate_and_score_completions( # Decode the generated completions completions_text = self.processing_class.batch_decode(completion_ids, skip_special_tokens=True) - # if is_conversational(inputs[0]): - # completions = [] - # for prompt, completion in zip(prompts, completions_text): - # bootstrap = prompt.pop()["content"] if prompt[-1]["role"] == "assistant" else "" - # completions.append([{"role": "assistant", "content": bootstrap + completion}]) - # else: - completions = completions_text rewards_per_func = torch.zeros(prompt_ids.size(0), len(self.reward_funcs), device=device) for i, (reward_func, reward_processing_class) in enumerate( zip(self.reward_funcs, self.reward_processing_classes) ): if isinstance(reward_func, nn.Module): # Module instead of PretrainedModel for compat with compiled models - # if is_conversational(inputs[0]): - # messages = [{"messages": p + c} for p, c in zip(prompts, completions)] - # texts = [apply_chat_template(x, reward_processing_class)["text"] for x in messages] - # else: - # texts = [p + c for p, c in zip(prompts, completions)] - texts = completions + texts = completions_text reward_inputs = reward_processing_class( texts, return_tensors="pt", padding=True, padding_side="right", add_special_tokens=False ) @@ -227,7 +183,7 @@ def _generate_and_score_completions( # Repeat all input columns (but "prompt" and "completion") to match the number of generations keys = [key for key in inputs[0] if key not in ["prompt", "completion"]] reward_kwargs = {key: [example[key] for example in inputs] for key in keys} - output_reward_func = reward_func(prompts=prompts, completions=completions, **reward_kwargs) + output_reward_func = reward_func(prompts=prompts, completions=completions_text, **reward_kwargs) rewards_per_func[:, i] = torch.tensor(output_reward_func, dtype=torch.float32, device=device) # Gather the reward per function: this part is crucial, because the rewards are normalized per group and the @@ -271,7 +227,6 @@ def _generate_and_score_completions( self._metrics[mode]["reward_std"].append(std_grouped_rewards.mean().item()) if self.log_completions and self.state.global_step % self.args.logging_steps == 0: - # prompts_to_log = gather_object(prompts_text) prompts_to_log = gather_object(prompts) completions_to_log = gather_object(completions_text) rewards_to_log = rewards.tolist() @@ -308,17 +263,17 @@ def _generate_and_score_completions( "ref_per_token_logps": ref_per_token_logps, "advantages": advantages, } + @profiling_decorator def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None): if return_outputs: raise ValueError("The GRPOTrainer does not support returning outputs") # Compute the per-token log probabilities for the model - prompt_ids = inputs['prompt_ids'] # torch.stack([x["input_ids"] for x in inputs]) - prompt_mask = inputs['prompt_mask'] # torch.stack([x["attention_mask"] for x in inputs]) - mel = inputs['mel'] # torch.stack([x["mel"] for x in inputs]) - mel_len = inputs['mel_len'] # torch.stack([torch.tensor(x["mel_len"]) for x in inputs]) + prompt_ids = inputs['prompt_ids'] + prompt_mask = inputs['prompt_mask'] + mel = inputs['mel'] + mel_len = inputs['mel_len'] - # prompt_ids, prompt_mask = inputs["prompt_ids"], inputs["prompt_mask"] completion_ids, completion_mask = inputs["completion_ids"], inputs["completion_mask"] input_ids = torch.cat([prompt_ids, completion_ids], dim=1) attention_mask = torch.cat([prompt_mask, completion_mask], dim=1) @@ -359,6 +314,7 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N self._metrics[mode]["clip_ratio"].append(self.accelerator.gather_for_metrics(clip_ratio).mean().item()) return loss + @profiling_decorator def _get_per_token_logps(self, model, input_ids, attention_mask, mel, mel_len, logits_to_keep): # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded logits = model(input_ids=input_ids, @@ -366,7 +322,7 @@ def _get_per_token_logps(self, model, input_ids, attention_mask, mel, mel_len, l mel=mel, mel_len=mel_len, logits_to_keep=logits_to_keep + 1).logits - # logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred + logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred input_ids = input_ids[:, -logits_to_keep:] # For transformers<=4.48, logits_to_keep argument isn't supported, so here we drop logits ourselves. diff --git a/speech_llm.py b/speech_llm.py index 9cd87ee..da36b18 100644 --- a/speech_llm.py +++ b/speech_llm.py @@ -2,7 +2,8 @@ import math from typing import Optional -from dataclasses import dataclass, field +import dataclasses +from dataclasses import dataclass, field, asdict from typing import Any, Callable, Optional, Sized, Union import safetensors @@ -10,13 +11,23 @@ import torch.nn as nn import torch.nn.functional as F import transformers -from transformers import AutoModelForCausalLM, PreTrainedModel +from transformers import AutoModelForCausalLM, PreTrainedModel, PretrainedConfig import wenet import whisper +from whisper.model import ModelDimensions, Whisper +import inspect @dataclass class ModelArguments: + + def __iter__(self): + for field in dataclasses.fields(self): + yield (field.name, getattr(self, field.name)) + for attr, value in inspect.getmembers(self.__class__): + if isinstance(value, property): + yield (attr,getattr(self, attr)) + llm_model_name_or_path: Optional[str] = field(default="Qwen/Qwen2-7B") whisper_model_name_or_path: Optional[str] = field(default="tiny") wenet_model_name_or_path: Optional[str] = field(default="") @@ -43,19 +54,49 @@ class ModelArguments: def ds_rate(self): return self.encoder_ds_rate * self.encoder_projector_ds_rate + @ds_rate.setter + def ds_rate(self, value): + if value != self.encoder_ds_rate * self.encoder_projector_ds_rate: + raise ValueError(" ds_rate != encoder_ds_rate / encoder_projector_ds_rate") + @property def speech_tokens_per_second(self): return self.frames_per_second / self.ds_rate + @speech_tokens_per_second.setter + def speech_tokens_per_second(self, value): + if value != self.frames_per_second / self.ds_rate: + raise ValueError(" speech_tokens_per_second != frames_per_second / ds_rate") + @property def max_speech_token_size(self): return math.ceil(self.max_speech_seconds * self.speech_tokens_per_second) + @max_speech_token_size.setter + def max_speech_token_size(self, value): + if value != math.ceil(self.max_speech_seconds * self.speech_tokens_per_second): + raise ValueError(" max_speech_token_size != max_speech_seconds * speech_tokens_per_second") + @property def max_mel_size(self): return self.max_speech_seconds * self.frames_per_second + @max_mel_size.setter + def max_mel_size(self, value): + if value != self.max_speech_seconds * self.frames_per_second: + raise ValueError(" max_mel_size != max_speech_seconds * frames_per_second") + +class WestSpeechConfig(PretrainedConfig): + model_type = "west_speech_model" + + def __init__(self, config:ModelArguments=None, **kwargs): + super().__init__(**kwargs) + self._name_or_path = "./west-slm" + if config is not None: + # config = ModelArguments() + for attr, value in config: + setattr(self, attr, value) def ctc_reduce(hyp, blank_id: int = 0): new_hyp = [] @@ -84,6 +125,7 @@ def __init__(self, config, encoder_dim, llm_dim): self.linear2 = nn.Linear(config.projector_hidden_size, llm_dim) self.relu2 = nn.ReLU() + @torch.autocast(device_type="cuda", dtype=torch.bfloat16) def forward(self, x): x = x.transpose(1, 2) x = self.conv1d(x) @@ -101,24 +143,79 @@ def freeze_model(model): class SpeechLLM(PreTrainedModel): + config_class = WestSpeechConfig supports_gradient_checkpointing = True def __init__( self, - llm: nn.Module, - encoder: nn.Module, - projector: nn.Module, config, - model_args: ModelArguments, + model_args: ModelArguments = None, + **kwargs ): super().__init__(config) - self.llm = llm + llm_config = transformers.AutoConfig.from_pretrained( + config.llm_model_name_or_path) + # llm_config.use_cache = False + + if 1: + # if model_args is not None: + if config.encoder_type == "whisper": + if model_args is not None: + encoder = whisper.load_model(config.whisper_model_name_or_path) + for field in dataclasses.fields(encoder.dims): + setattr(config, f'whisper_{field.name}', getattr(encoder.dims, field.name)) + else: + whisper_dim = ModelDimensions(n_mels=1, + n_audio_ctx=1 , + n_audio_state=1 , + n_audio_head=1 , + n_audio_layer=1 , + n_vocab=1 , + n_text_ctx=1 , + n_text_state=1 , + n_text_head=1 , + n_text_layer=1) + for field in dataclasses.fields(whisper_dim): + v = getattr(config, f'whisper_{field.name}') + setattr(whisper_dim, field.name, v) + encoder = Whisper(whisper_dim) + + + elif config.encoder_type == "wenet": + encoder = wenet.load_model_pt(config.wenet_model_name_or_path) + device = "cuda" if torch.cuda.is_available() else "cpu" + encoder = encoder.to(device) + else: + raise ValueError(f"Unexpected encoder type {config.encoder_type}") + + # Load llm model and tokenizer + llm_model = AutoModelForCausalLM.from_pretrained( + config.llm_model_name_or_path, + config=llm_config, + torch_dtype='auto', + ) + if config.encoder_type == "whisper": + encoder_dim = encoder.dims.n_audio_state + else: + encoder_dim = encoder.encoder.output_size() + + config.encoder_dim = encoder_dim + if config.encoder_dim: + encoder_dim = config.encoder_dim + + config.hidden_size = llm_config.hidden_size + llm_dim = llm_config.hidden_size + projector = ProjectorCov1d(config, encoder_dim, llm_dim) + total_params = sum(p.numel() for p in projector.parameters()) + print('Projector total params: {:.2f}M'.format(total_params / 1024 / 1024)) + + self.llm = llm_model self.encoder = encoder self.projector = projector - self._keys_to_ignore_on_save = set() + # self._keys_to_ignore_on_save = set() # Do not save the parameter of llm and whisper - for k in self.llm.state_dict().keys(): - self._keys_to_ignore_on_save.add('llm.' + k) + # for k in self.llm.state_dict().keys(): + # self._keys_to_ignore_on_save.add('llm.' + k) # for k in self.encoder.state_dict().keys(): # self._keys_to_ignore_on_save.add('encoder.' + k) # Use bos_token_id as CTC blank id @@ -126,7 +223,7 @@ def __init__( reduction='mean', zero_infinity=True) self.blank_id = config.bos_token_id - self.model_args = model_args + self.model_args = config def get_speech_embeddings(self, mel, mel_len): max_speech_size = self.model_args.max_speech_token_size @@ -188,11 +285,7 @@ def generate( attention_mask: Optional[torch.Tensor] = None, mel: torch.LongTensor = None, mel_len: torch.LongTensor = None, - eos_token_id=None, decode_config=None, - do_sample=False, - top_p=1.0, - temperature=0.7, **kwargs ): max_speech_size = self.model_args.max_speech_token_size @@ -203,12 +296,12 @@ def generate( model_outputs = self.llm.generate( inputs_embeds=inputs_embeds, attention_mask=attention_mask, - do_sample=do_sample, - top_p=top_p, - temperature=temperature, + do_sample=decode_config.do_sample, + top_p=decode_config.top_p, + temperature=decode_config.temperature, num_beams=decode_config.num_beams, max_new_tokens=decode_config.max_new_tokens, - eos_token_id=eos_token_id, + eos_token_id=decode_config.eos_token_id, pad_token_id=decode_config.pad_token_id, **kwargs ) @@ -241,6 +334,10 @@ def decode_ctc( def enable_input_require_grads(self): self.llm.enable_input_require_grads() + def freeze_projector(self): + freeze_model(self.projector) + # self.projector.eval() + def freeze_encoder(self): freeze_model(self.encoder) self.encoder.eval() @@ -254,33 +351,8 @@ def load_projector(self, projector_path): def init_model(model_args): - if model_args.encoder_type == "whisper": - encoder = whisper.load_model(model_args.whisper_model_name_or_path) - elif model_args.encoder_type == "wenet": - encoder = wenet.load_model_pt(model_args.wenet_model_name_or_path) - device = "cuda" if torch.cuda.is_available() else "cpu" - encoder = encoder.to(device) - else: - raise ValueError(f"Unexpected encoder type {model_args.encoder_type}") - - # Load llm model and tokenizer - config = transformers.AutoConfig.from_pretrained( - model_args.llm_model_name_or_path) - config.use_cache = False - llm_model = AutoModelForCausalLM.from_pretrained( - model_args.llm_model_name_or_path, - config=config, - torch_dtype='auto', - ) - if model_args.encoder_type == "whisper": - encoder_dim = encoder.dims.n_audio_state - else: - encoder_dim = encoder.encoder.output_size() - llm_dim = config.hidden_size - projector = ProjectorCov1d(model_args, encoder_dim, llm_dim) - total_params = sum(p.numel() for p in projector.parameters()) - print('Projector total params: {:.2f}M'.format(total_params / 1024 / 1024)) - model = SpeechLLM(llm_model, encoder, projector, config, model_args) + west_model_config = WestSpeechConfig(config=model_args) + model = SpeechLLM(west_model_config, model_args=model_args) if model_args.projector_model_path is not None: - model.load_projector(model_args.projector_model_path) + model.load_projector(model_args.projector_model_path,) return model diff --git a/train.py b/train.py index 8febb25..95213fe 100644 --- a/train.py +++ b/train.py @@ -9,7 +9,7 @@ from transformers import AutoTokenizer, Trainer from dataset import DataArguments, SpeechDataset -from speech_llm import init_model, ModelArguments +from speech_llm import init_model, ModelArguments, SpeechLLM from trl import GRPOConfig from speech_grpo_trainer import SpeechGRPOTrainer from reward_funcs import active_reward_func @@ -19,6 +19,7 @@ class TrainingArguments(transformers.TrainingArguments): optim: str = field(default="adafactor") grpo: bool = field(default=False) remove_unused_columns: bool = field(default=False) + temperature: float = field(default=0.9) def main(): @@ -32,7 +33,8 @@ def main(): model = init_model(model_args) model.freeze_llm() - model.freeze_encoder() + # model.freeze_encoder() + if training_args.gradient_checkpointing: model.enable_input_require_grads() @@ -65,6 +67,11 @@ def main(): 'reward_funcs': active_reward_func } args = grpo_config + ref_model = SpeechLLM.from_pretrained('./west-slm') + ref_model.freeze_llm() + ref_model.freeze_encoder() + ref_model.freeze_projector() + ref_model.eval() else: trainer_cls = Trainer trainer_arg = { @@ -75,7 +82,7 @@ def main(): print(args.to_dict()) # Start trainer - trainer = trainer_cls(model=model, + trainer = trainer_cls(model=model, ref_model=ref_model, args=args, train_dataset=train_dataset, eval_dataset=eval_dataset, From 797919cc8414e6eca21224770cade64c15a63b8c Mon Sep 17 00:00:00 2001 From: libo Date: Fri, 28 Feb 2025 10:03:53 +0000 Subject: [PATCH 4/4] fix reward --- reward_funcs.py | 3 ++- run.sh | 3 ++- train.py | 3 +-- 3 files changed, 5 insertions(+), 4 deletions(-) diff --git a/reward_funcs.py b/reward_funcs.py index a9f270e..a67f3c5 100644 --- a/reward_funcs.py +++ b/reward_funcs.py @@ -11,7 +11,8 @@ def editdistance_score(completions, **kwargs): diff = [] for hyp, lab in zip(completions,kwargs['txt']): # diff.append(-1.0 * editdistance.eval(hyp, lab) / len(lab)) - diff.append(np.log(1e-9 + editdistance.eval(hyp, lab))) + # diff.append(np.log(1e-9 + editdistance.eval(hyp, lab))) + diff.append(-editdistance.eval(hyp, lab)) return diff def word_count(completions, **kwargs): diff --git a/run.sh b/run.sh index eaf818e..f761a34 100755 --- a/run.sh +++ b/run.sh @@ -74,8 +74,9 @@ if [ ${stage} -le 1 ] && [ ${stop_stage} -ge 1 ]; then fi +ckpt=$2 TEST_MODEL=Qwen-1.5B-Instruct-whisper-tiny/checkpoint-1170/model.safetensors -TEST_MODEL=Qwen/Qwen2-1.5B-Instruct-whisper-tiny/checkpoint-3/model.safetensors +TEST_MODEL=Qwen/Qwen2-1.5B-Instruct-whisper-tiny/checkpoint-${ckpt}/model.safetensors if [ ${stage} -le 2 ] && [ ${stop_stage} -ge 2 ]; then python recognize.py \ --llm_model_name_or_path ${LLM_MODEL} \ diff --git a/train.py b/train.py index 95213fe..ae505dd 100644 --- a/train.py +++ b/train.py @@ -33,8 +33,7 @@ def main(): model = init_model(model_args) model.freeze_llm() - # model.freeze_encoder() - + model.freeze_encoder() if training_args.gradient_checkpointing: model.enable_input_require_grads()