From 7eae14f3cee8118ac488286d5e034f9e9c9eb29d Mon Sep 17 00:00:00 2001 From: jingyaogong Date: Sat, 27 Dec 2025 07:14:36 +0800 Subject: [PATCH] [feat] remove empty_cache --- trainer/train_distill_reason.py | 3 +-- trainer/train_distillation.py | 3 +-- trainer/train_dpo.py | 3 +-- trainer/train_full_sft.py | 5 ++--- trainer/train_grpo.py | 5 +---- trainer/train_lora.py | 3 +-- trainer/train_ppo.py | 4 +--- trainer/train_pretrain.py | 5 ++--- trainer/train_spo.py | 5 +---- trainer/trainer_utils.py | 1 - 10 files changed, 11 insertions(+), 26 deletions(-) diff --git a/trainer/train_distill_reason.py b/trainer/train_distill_reason.py index e86fe9d..be69738 100644 --- a/trainer/train_distill_reason.py +++ b/trainer/train_distill_reason.py @@ -65,7 +65,6 @@ def train_epoch(epoch, loader, iters, tokenizer, lm_config, start_step=0, wandb= scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_none=True) - torch.cuda.empty_cache() if step % args.log_interval == 0 or step == iters - 1: spend_time = time.time() - start_time @@ -103,7 +102,7 @@ if __name__ == "__main__": parser.add_argument("--learning_rate", type=float, default=1e-6, help="初始学习率") parser.add_argument("--device", type=str, default="cuda:0" if torch.cuda.is_available() else "cpu", help="训练设备") parser.add_argument("--dtype", type=str, default="bfloat16", help="混合精度类型") - parser.add_argument("--num_workers", type=int, default=1, help="数据加载线程数") + parser.add_argument("--num_workers", type=int, default=8, help="数据加载线程数") parser.add_argument("--accumulation_steps", type=int, default=1, help="梯度累积步数") parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值") parser.add_argument("--log_interval", type=int, default=100, help="日志打印间隔") diff --git a/trainer/train_distillation.py b/trainer/train_distillation.py index e4b4ecd..712bac6 100644 --- a/trainer/train_distillation.py +++ b/trainer/train_distillation.py @@ -96,7 +96,6 @@ def train_epoch(epoch, loader, iters, teacher_model, lm_config_student, start_st scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_none=True) - torch.cuda.empty_cache() if step % args.log_interval == 0 or step == iters - 1: spend_time = time.time() - start_time @@ -141,7 +140,7 @@ if __name__ == "__main__": parser.add_argument("--learning_rate", type=float, default=5e-6, help="初始学习率") parser.add_argument("--device", type=str, default="cuda:0" if torch.cuda.is_available() else "cpu", help="训练设备") parser.add_argument("--dtype", type=str, default="bfloat16", help="混合精度类型") - parser.add_argument("--num_workers", type=int, default=1, help="数据加载线程数") + parser.add_argument("--num_workers", type=int, default=8, help="数据加载线程数") parser.add_argument("--accumulation_steps", type=int, default=1, help="梯度累积步数") parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值") parser.add_argument("--log_interval", type=int, default=100, help="日志打印间隔") diff --git a/trainer/train_dpo.py b/trainer/train_dpo.py index a55a450..b20b53f 100644 --- a/trainer/train_dpo.py +++ b/trainer/train_dpo.py @@ -90,7 +90,6 @@ def train_epoch(epoch, loader, iters, ref_model, lm_config, start_step=0, wandb= scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_none=True) - torch.cuda.empty_cache() if step % args.log_interval == 0 or step == iters - 1: spend_time = time.time() - start_time @@ -129,7 +128,7 @@ if __name__ == "__main__": parser.add_argument("--learning_rate", type=float, default=4e-8, help="初始学习率(建议<=5e-8避免遗忘)") parser.add_argument("--device", type=str, default="cuda:0" if torch.cuda.is_available() else "cpu", help="训练设备") parser.add_argument("--dtype", type=str, default="bfloat16", help="混合精度类型") - parser.add_argument("--num_workers", type=int, default=1, help="数据加载线程数") + parser.add_argument("--num_workers", type=int, default=8, help="数据加载线程数") parser.add_argument("--accumulation_steps", type=int, default=1, help="梯度累积步数") parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值") parser.add_argument("--log_interval", type=int, default=100, help="日志打印间隔") diff --git a/trainer/train_full_sft.py b/trainer/train_full_sft.py index f0489fb..9b7f011 100644 --- a/trainer/train_full_sft.py +++ b/trainer/train_full_sft.py @@ -52,7 +52,6 @@ def train_epoch(epoch, loader, iters, start_step=0, wandb=None): scaler.update() optimizer.zero_grad(set_to_none=True) - torch.cuda.empty_cache() if step % args.log_interval == 0 or step == iters - 1: spend_time = time.time() - start_time @@ -91,11 +90,11 @@ if __name__ == "__main__": parser.add_argument("--learning_rate", type=float, default=5e-7, help="初始学习率") parser.add_argument("--device", type=str, default="cuda:0" if torch.cuda.is_available() else "cpu", help="训练设备") parser.add_argument("--dtype", type=str, default="bfloat16", help="混合精度类型") - parser.add_argument("--num_workers", type=int, default=1, help="数据加载线程数") + parser.add_argument("--num_workers", type=int, default=8, help="数据加载线程数") parser.add_argument("--accumulation_steps", type=int, default=1, help="梯度累积步数") parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值") parser.add_argument("--log_interval", type=int, default=100, help="日志打印间隔") - parser.add_argument("--save_interval", type=int, default=100, help="模型保存间隔") + parser.add_argument("--save_interval", type=int, default=1000, help="模型保存间隔") parser.add_argument('--hidden_size', default=512, type=int, help="隐藏层维度") parser.add_argument('--num_hidden_layers', default=8, type=int, help="隐藏层数量") parser.add_argument('--max_seq_len', default=340, type=int, help="训练的最大截断长度(中文1token≈1.5~1.7字符)") diff --git a/trainer/train_grpo.py b/trainer/train_grpo.py index 4e11f22..897d9a8 100755 --- a/trainer/train_grpo.py +++ b/trainer/train_grpo.py @@ -149,7 +149,6 @@ def grpo_train_epoch(epoch, loader, iters, ref_model, reward_model, reward_token optimizer.step() scheduler.step() optimizer.zero_grad() - torch.cuda.empty_cache() if step % args.log_interval == 0 or step == iters: policy_loss_val = loss.item() @@ -183,8 +182,6 @@ def grpo_train_epoch(epoch, loader, iters, ref_model, reward_model, reward_token del prompt_inputs, outputs, completion_ids, per_token_logps, ref_per_token_logps del completions, rewards, grouped_rewards, mean_r, std_r, advantages, completion_mask - torch.cuda.empty_cache() - gc.collect() if __name__ == "__main__": @@ -196,7 +193,7 @@ if __name__ == "__main__": parser.add_argument("--learning_rate", type=float, default=8e-8, help="初始学习率") parser.add_argument("--device", type=str, default="cuda:0" if torch.cuda.is_available() else "cpu", help="训练设备") parser.add_argument("--dtype", type=str, default="bfloat16", help="混合精度类型") - parser.add_argument("--num_workers", type=int, default=1, help="数据加载线程数") + parser.add_argument("--num_workers", type=int, default=8, help="数据加载线程数") parser.add_argument("--accumulation_steps", type=int, default=1, help="梯度累积步数") parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值") parser.add_argument("--log_interval", type=int, default=1, help="日志打印间隔") diff --git a/trainer/train_lora.py b/trainer/train_lora.py index 474ae18..89cb7a9 100644 --- a/trainer/train_lora.py +++ b/trainer/train_lora.py @@ -53,7 +53,6 @@ def train_epoch(epoch, loader, iters, lora_params, start_step=0, wandb=None): scaler.update() optimizer.zero_grad(set_to_none=True) - torch.cuda.empty_cache() if step % args.log_interval == 0 or step == iters - 1: spend_time = time.time() - start_time @@ -85,7 +84,7 @@ if __name__ == "__main__": parser.add_argument("--learning_rate", type=float, default=1e-4, help="初始学习率") parser.add_argument("--device", type=str, default="cuda:0" if torch.cuda.is_available() else "cpu", help="训练设备") parser.add_argument("--dtype", type=str, default="bfloat16", help="混合精度类型") - parser.add_argument("--num_workers", type=int, default=1, help="数据加载线程数") + parser.add_argument("--num_workers", type=int, default=8, help="数据加载线程数") parser.add_argument("--accumulation_steps", type=int, default=1, help="梯度累积步数") parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值") parser.add_argument("--log_interval", type=int, default=10, help="日志打印间隔") diff --git a/trainer/train_ppo.py b/trainer/train_ppo.py index 836bf67..cb0ec38 100644 --- a/trainer/train_ppo.py +++ b/trainer/train_ppo.py @@ -179,7 +179,6 @@ def ppo_train_epoch(epoch, loader, iters, old_actor_model, ref_model, actor_sche critic_scheduler.step() actor_optimizer.zero_grad() critic_optimizer.zero_grad() - torch.cuda.empty_cache() if is_main_process(): response_ids = gen_out[:, enc.input_ids.shape[1]:] @@ -237,7 +236,6 @@ def ppo_train_epoch(epoch, loader, iters, old_actor_model, ref_model, actor_sche del enc, gen_out, responses_text, rewards, full_mask, values_seq, values, advantages del logits, labels, logp_tokens, final_mask, actor_logp, old_logits, old_logp, ref_logits, ref_logp del kl, kl_ref, ratio, surr1, surr2, policy_loss, value_loss, loss - torch.cuda.empty_cache() if __name__ == "__main__": @@ -250,7 +248,7 @@ if __name__ == "__main__": parser.add_argument("--critic_learning_rate", type=float, default=8e-8, help="Critic学习率") parser.add_argument("--device", type=str, default="cuda:0" if torch.cuda.is_available() else "cpu", help="训练设备") parser.add_argument("--dtype", type=str, default="bfloat16", help="混合精度类型") - parser.add_argument("--num_workers", type=int, default=1, help="数据加载线程数") + parser.add_argument("--num_workers", type=int, default=8, help="数据加载线程数") parser.add_argument("--accumulation_steps", type=int, default=1, help="梯度累积步数") parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值") parser.add_argument("--log_interval", type=int, default=1, help="日志打印间隔") diff --git a/trainer/train_pretrain.py b/trainer/train_pretrain.py index d02d9b5..5f05341 100644 --- a/trainer/train_pretrain.py +++ b/trainer/train_pretrain.py @@ -52,7 +52,6 @@ def train_epoch(epoch, loader, iters, start_step=0, wandb=None): scaler.update() optimizer.zero_grad(set_to_none=True) - torch.cuda.empty_cache() if step % args.log_interval == 0 or step == iters - 1: spend_time = time.time() - start_time @@ -90,11 +89,11 @@ if __name__ == "__main__": parser.add_argument("--learning_rate", type=float, default=5e-4, help="初始学习率") parser.add_argument("--device", type=str, default="cuda:0" if torch.cuda.is_available() else "cpu", help="训练设备") parser.add_argument("--dtype", type=str, default="bfloat16", help="混合精度类型") - parser.add_argument("--num_workers", type=int, default=1, help="数据加载线程数") + parser.add_argument("--num_workers", type=int, default=8, help="数据加载线程数") parser.add_argument("--accumulation_steps", type=int, default=8, help="梯度累积步数") parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值") parser.add_argument("--log_interval", type=int, default=100, help="日志打印间隔") - parser.add_argument("--save_interval", type=int, default=100, help="模型保存间隔") + parser.add_argument("--save_interval", type=int, default=1000, help="模型保存间隔") parser.add_argument('--hidden_size', default=512, type=int, help="隐藏层维度") parser.add_argument('--num_hidden_layers', default=8, type=int, help="隐藏层数量") parser.add_argument('--max_seq_len', default=340, type=int, help="训练的最大截断长度(中文1token≈1.5~1.7字符)") diff --git a/trainer/train_spo.py b/trainer/train_spo.py index bac7c14..37493e4 100755 --- a/trainer/train_spo.py +++ b/trainer/train_spo.py @@ -192,7 +192,6 @@ def spo_train_epoch(epoch, loader, iters, ref_model, reward_model, reward_tokeni optimizer.step() scheduler.step() optimizer.zero_grad() - torch.cuda.empty_cache() if step % args.log_interval == 0 or step == iters: policy_loss_val = loss.item() @@ -231,8 +230,6 @@ def spo_train_epoch(epoch, loader, iters, ref_model, reward_model, reward_tokeni del prompt_inputs, outputs, completion_ids, per_token_logps, ref_per_token_logps del completions, rewards, advantages, completion_mask, baselines, response_masks - torch.cuda.empty_cache() - gc.collect() if __name__ == "__main__": @@ -244,7 +241,7 @@ if __name__ == "__main__": parser.add_argument("--learning_rate", type=float, default=1e-7, help="初始学习率") parser.add_argument("--device", type=str, default="cuda:0" if torch.cuda.is_available() else "cpu", help="训练设备") parser.add_argument("--dtype", type=str, default="bfloat16", help="混合精度类型") - parser.add_argument("--num_workers", type=int, default=1, help="数据加载线程数") + parser.add_argument("--num_workers", type=int, default=8, help="数据加载线程数") parser.add_argument("--accumulation_steps", type=int, default=4, help="梯度累积步数") parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值") parser.add_argument("--log_interval", type=int, default=1, help="日志打印间隔") diff --git a/trainer/trainer_utils.py b/trainer/trainer_utils.py index c1a4ca2..e2dd433 100644 --- a/trainer/trainer_utils.py +++ b/trainer/trainer_utils.py @@ -88,7 +88,6 @@ def lm_checkpoint(lm_config, weight='full_sft', model=None, optimizer=None, epoc torch.save(resume_data, resume_tmp) os.replace(resume_tmp, resume_path) del state_dict, resume_data - gc.collect() torch.cuda.empty_cache() else: # 加载模式 if os.path.exists(resume_path):