You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

梯度累加训练仅完成50%即终止的问题求助

梯度累加训练异常终止问题

训练进度停在50%

  • 原batch_size=16,设置accumulation=2以近似batch_size=32的训练效果
  • 原训练耗时1小时,预期开启梯度累加后时长为2小时
  • 实际训练仅完成50%就终止,耗时仍为1小时

问题代码

def train_runner(model, train_dataset, valid_dataset , batch_size, num_train_epochs, learning_rate):
device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')

    model.to(device)
    model.train()
    train_dataloader = DataLoader(dataset=train_dataset, batch_size=batch_size)
    valid_dataloader = DataLoader(dataset = valid_dataset, batch_size = batch_size)

    lowest_total_valid_loss = 9999.
    step = 0
    global_total_step = len(train_dataloader) * num_train_epochs
    optimizer = AdamW(model.parameters(), lr=learning_rate, weight_decay=0)
    print("TRAIN START")
    with tqdm(total=global_total_step, unit='step') as t:
        total = 0
        total_loss = 0
        for epoch in range(num_train_epochs):
            for iteration,batch in enumerate(train_dataloader):
                #optimizer.zero_grad()
                input_ids = batch['input_ids'].to(device)
                attention_mask = batch['attention_mask'].to(device)
                start_positions = batch['start_positions'].to(device)
                end_positions = batch['end_positions'].to(device)
                outputs = model(input_ids,
                             attention_mask=attention_mask,
                             start_positions=start_positions,
                             end_positions=end_positions)
                loss = outputs.loss
                (loss / ACCUMULATION).backward()

                step += 1
                if step % ACCUMULATION:
                    continue

                clip_grad_norm_(model.parameters(), max_norm=1.)
                optimizer.step()
                optimizer.zero_grad(set_to_none=True)
                
                batch_loss = loss.item() * len(input_ids)
                total += len(input_ids)
                total_loss += batch_loss / ACCUMULATION
                global_total_step += 1
                t.set_postfix(loss="{:.6f}".format(total_loss / total), batch_loss="{:.6f}".format(batch_loss))
                t.update(1)
                
                del input_ids
                del attention_mask
                del start_positions
                del end_positions
                del outputs
                del loss

                ## validation ##
                if iteration != 0 and iteration % int(len(train_dataloader) / 10) == 0:
                    total_valid_loss = 0
                    for batch_val in valid_dataloader:
                        model.eval()
                        optimizer.zero_grad()

                        input_ids = batch_val['input_ids'].to(device)
                        attention_mask = batch_val['attention_mask'].to(device)
                        start_positions = batch_val['start_positions'].to(device)
                        end_positions = batch_val['end_positions'].to(device)
                
                        with torch.no_grad():
                            outputs = model(input_ids,
                                    attention_mask=attention_mask,
                                    start_positions=start_positions,
                                    end_positions=end_positions)
                            loss = outputs.loss
                            total_valid_loss += loss.item()
                    
                    if total_valid_loss < lowest_total_valid_loss:
                        print(f"lowest_total_valid_loss: {total_valid_loss} epoch : {epoch} iteration : {iteration}")
                        torch.save(model.state_dict(),'./output_model_best')
                        lowest_total_valid_loss = total_valid_loss
                ## validation ##

#model.save_pretrained("./klue_output_model")
print("TRAIN END")

问题根源分析

  1. tqdm总步数逻辑混乱
    初始global_total_step是单步更新的总步数,但梯度累加后实际参数更新步数应为len(train_dataloader) * num_train_epochs // ACCUMULATION。更关键的是你在循环内每次更新参数时都执行global_total_step += 1,导致tqdm初始化的总步数和实际运行逻辑冲突,当实际更新步数达到初始总步数时,tqdm判定任务完成,直接终止循环。

  2. 梯度累加判断逻辑完全反转
    if step % ACCUMULATION: 等价于if step % ACCUMULATION != 0,这会导致只有当步数不是累加倍数时才跳过更新,和你需要的“达到累加步数再更新”逻辑完全相反。

修正后的代码

def train_runner(model, train_dataset, valid_dataset , batch_size, num_train_epochs, learning_rate):
    ACCUMULATION = 2  # 显式定义累加系数,建议作为参数传入
    device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')

    model.to(device)
    model.train()
    train_dataloader = DataLoader(dataset=train_dataset, batch_size=batch_size)
    valid_dataloader = DataLoader(dataset = valid_dataset, batch_size = batch_size)

    lowest_total_valid_loss = 9999.
    step = 0
    # 修正总步数:参数更新总次数 = 总batch数 / 累加次数 * epoch数
    global_total_step = (len(train_dataloader) // ACCUMULATION) * num_train_epochs
    optimizer = AdamW(model.parameters(), lr=learning_rate, weight_decay=0)
    print("TRAIN START")
    with tqdm(total=global_total_step, unit='step') as t:
        total = 0
        total_loss = 0
        for epoch in range(num_train_epochs):
            for iteration,batch in enumerate(train_dataloader):
                input_ids = batch['input_ids'].to(device)
                attention_mask = batch['attention_mask'].to(device)
                start_positions = batch['start_positions'].to(device)
                end_positions = batch['end_positions'].to(device)
                outputs = model(input_ids,
                             attention_mask=attention_mask,
                             start_positions=start_positions,
                             end_positions=end_positions)
                loss = outputs.loss
                (loss / ACCUMULATION).backward()

                step += 1
                # 修正判断:未达到累加步数时跳过参数更新
                if step % ACCUMULATION != 0:
                    del input_ids, attention_mask, start_positions, end_positions, outputs, loss
                    continue

                clip_grad_norm_(model.parameters(), max_norm=1.)
                optimizer.step()
                optimizer.zero_grad(set_to_none=True)
                
                batch_loss = loss.item() * len(input_ids)
                total += len(input_ids) * ACCUMULATION  # 对应实际有效batch大小
                total_loss += batch_loss
                t.set_postfix(loss="{:.6f}".format(total_loss / total), batch_loss="{:.6f}".format(batch_loss))
                t.update(1)
                
                del input_ids, attention_mask, start_positions, end_positions, outputs, loss

                ## validation ##
                if iteration != 0 and iteration % int(len(train_dataloader) / 10) == 0:
                    total_valid_loss = 0
                    model.eval()
                    for batch_val in valid_dataloader:
                        input_ids = batch_val['input_ids'].to(device)
                        attention_mask = batch_val['attention_mask'].to(device)
                        start_positions = batch_val['start_positions'].to(device)
                        end_positions = batch_val['end_positions'].to(device)
                
                        with torch.no_grad():
                            outputs = model(input_ids,
                                    attention_mask=attention_mask,
                                    start_positions=start_positions,
                                    end_positions=end_positions)
                            loss = outputs.loss
                            total_valid_loss += loss.item()
                    
                    if total_valid_loss < lowest_total_valid_loss:
                        print(f"lowest_total_valid_loss: {total_valid_loss} epoch : {epoch} iteration : {iteration}")
                        torch.save(model.state_dict(),'./output_model_best')
                        lowest_total_valid_loss = total_valid_loss
                    model.train()  # 验证后切回训练模式
                ## validation ##

    #model.save_pretrained("./klue_output_model")
    print("TRAIN END")

额外注意事项

  • 验证完成后必须切回model.train(),否则后续训练会保持eval模式,无法计算梯度
  • 所有分支都要执行显存清理,避免未更新参数时的显存泄漏
  • loss统计要对应累加后的实际有效batch大小,保证显示的loss值准确

内容的提问来源于stack exchange,提问作者Justin

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.05 11:20:54