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

LSTM模型中Batch Size对精度的影响及相关技术疑问

LSTM模型Batch Size对性能影响的技术分析

问题背景

我基于PyTorch开发了一款LSTM模型,用于识别N变量序列中对应优良(1类)与不良(0类)结果的模式。使用含12个训练样本的测试用例验证时,发现batch size=1(单样本计算损失后立即更新参数)的模型区分效果远优于全量样本作为batch size(一次性计算所有样本损失再更新参数)的情况,具体结果如下:

实验结果

Batch Size=1时的最终预测结果

Train set targets: [[0.0]], predictions:[0.24779687821865082]
Train set targets: [[1.0]], predictions:[0.9567258954048157]
Train set targets: [[1.0]], predictions:[0.8191764950752258]
Train set targets: [[0.0]], predictions:[0.20435290038585663]
Train set targets: [[0.0]], predictions:[0.1295892596244812]
Train set targets: [[1.0]], predictions:[0.9186112284660339]
Train set targets: [[1.0]], predictions:[0.6797895431518555]
Train set targets: [[1.0]], predictions:[0.9642216563224792]
Train set targets: [[1.0]], predictions:[0.9764360785484314]
Train set targets: [[0.0]], predictions:[0.670409619808197]
Train set targets: [[1.0]], predictions:[0.7026165723800659]
Train set targets: [[1.0]], predictions:[0.8404821157455444]

Test set targets: [[0.0]], predictions:[0.08575308322906494]
Test set targets: [[1.0]], predictions:[0.7602054476737976]
Test set targets: [[1.0]], predictions:[0.7767713069915771]

全量样本作为Batch Size时的最终预测结果

Train set 
targets: [[0.0], [1.0], [0.0], [1.0], [1.0], [1.0], [1.0], [0.0], [1.0], [1.0], [0.0], [1.0]], 
predictions:[0.6307732462882996, 0.6873687505722046, 0.6007956862449646, 0.7481836080551147, 0.7676156759262085, 0.6568607091903687, 0.7259970307350159, 0.597843587398529, 0.6819412708282471, 0.6660482287406921, 0.5785030126571655, 0.7716434597969055]


Test set 
Targets: [[0.0], [1.0], [1.0]], 
Predictions:[0.36408719420433044, 0.7265898585319519, 0.7854364514350891]

代码实现

模型定义

class LSTMModel(nn.Module):
    def __init__(self, input_dim, hidden_dim):
        super(LSTMFrzModel, self).__init__()
        self.input_dim = input_dim
        self.hidden_dim = hidden_dim
        
        # LSTM
        self.lstm = nn.LSTM(input_dim, hidden_dim, batch_first=True)
        
        # Readout layer
        self.fc = nn.Linear(hidden_dim, 1)
        self.sigmoid = nn.Sigmoid()

    
    def forward(self, inputs):
        mlp_out, (hidden, _) = self.lstm(inputs)
        output = self.fc(hidden)
        output = self.sigmoid(output)
        return output

数据集代码

class LSTMDataset(Dataset):
  def __init__(self, x, y):
    self.x = x
    self.y = y

  def __len__(self):
    return len(self.y)
  
  def __getitem__(self,idx):
    inputs = [torch.from_numpy(self.x[idx][ts]).unsqueeze(0) for ts in range(len(self.x[idx]))]
    inputs = torch.cat(inputs)
    target = torch.Tensor([self.y[idx]])
    return inputs, target

训练代码

# This function is used to train a single model
def train_lstm(
    X_training,
    y_training,
    X_testing,
    y_testing,
    hidden_dim=10,
    lr=1e-4,
    f1_thresh = 0.5,
    use_gpu=False,
    batch_size=100,
    num_epochs=4000,
    assym_wt=0, #Equal weights for 0 and 1
):
    start_time = time.time()
    clf = LSTMModel(len(X_training[0][0]), hidden_dim)
    # Move to GPU if available
    use_gpu = use_gpu and torch.cuda.is_available()
    device = torch.device("cuda" if use_gpu else "cpu")

    # Define the loss function and optimizer
    optimizer = torch.optim.Adam(clf.parameters(), lr=lr)
    clf = clf.to(device)
    loss_function = nn.BCELoss()
    loss_function = loss_function.to(device)

    dataset = LSTMDataset(X_training, y_training)
    trainloader = torch.utils.data.DataLoader(dataset, batch_size=batch_size, shuffle=True)

    # Run the training loop
    # per_epoch_precision = []
    # per_epoch_recall = []
    train_preds = []
    train_targets = []
    for epoch in range(0, num_epochs):
        # Set current loss value
        current_loss = 0.0
        
        # Iterate over the DataLoader for training data
        clf.train()  # set to train mode
        for i, data in enumerate(trainloader):
            # Get inputs
            inputs, targets = data                

            # Zero the gradients
            optimizer.zero_grad()

            # Perform forward pass
            outputs = clf(inputs)
            # Store predictions/targets in last epoch to compute accuracy stats 
            if epoch == num_epochs - 1:
                train_targets += targets
                train_preds += outputs.view(-1).tolist()
                print(f'Train set targets: {targets.tolist()}, predictions:{outputs.view(-1).tolist()}')

            # Compute loss
            targets = torch.FloatTensor(targets).unsqueeze(0)
            # Apply assymetric weights to handle unbalanced datasets
            if assym_wt > 0:
                loss_function = nn.BCELoss(weight=assym_wt * targets + 1)
                loss_function = loss_function.to(device)
            loss = loss_function(outputs, targets)
            
            # Perform backward pass
            loss.backward()

            # Perform optimization
            optimizer.step()

            # Print statistics
            current_loss += loss.item()

        if (epoch % 250) == 249:
            print("Loss after epoch %5d: %.3f" % (epoch + 1, current_loss / 500))
            current_loss = 0.0

    # Process is complete.
    print("Training process has finished.")
    train_preds = torch.FloatTensor(train_preds)
    train_targets = torch.FloatTensor(train_targets)

    clf.eval()  # set to eval mode
    dataset = LSTMDataset(X_testing, y_testing)
    testloader = torch.utils.data.DataLoader(dataset, batch_size=batch_size, shuffle=True)
    test_preds = []
    test_targets = []
    with torch.no_grad():
        for i, data in enumerate(testloader):
            # Get inputs
            inputs, targets = data
            # Perform forward pass
            preds = clf(inputs)
            print(f'Test set targets: {targets.tolist()}, predictions:{preds.view(-1).tolist()}')
            # Store predictions and targets so that we can compute stats
            test_preds += preds.view(-1).tolist()
            test_targets += targets
            
        
        test_preds = torch.FloatTensor(test_preds)
        test_targets = torch.FloatTensor(test_targets)
        pr_fig = go.Figure()
        pr_fig.update_xaxes(title_text="Recall")
        pr_fig.update_yaxes(title_text="Precision")
        plot_pr_withtorch(test_targets, test_preds, pr_fig, "Test PR ")
        plot_pr_withtorch(train_targets, train_preds, pr_fig, "Train PR ")
        pr_fig.show()

        train_F1 = get_scores(train_targets, train_preds > f1_thresh, "Train set scores")
        test_F1 = get_scores(test_targets, test_preds > f1_thresh, "Test set scores")
        print(f"Train time is {time.time() - start_time}")

    return train_F1, test_F1

测试用例

# test case
train = np.array([[10, 20, 5, 10, 1], [10, 5, 8, 3, 0], [10, 5, 5, 10, 1]])
for ind in range(0, 3):
    for spread in [5, 10, 50]:
        newrow = spread + train[ind]
        newrow[-1] -= spread
        # print(newrow)
        train = np.vstack([train, newrow])

test = np.array([[0, 8, 3, 7, 1], [19, 8, 12, 3, 0], [1000, 450, 75, 135, 1]])
train_df = pd.DataFrame(train, columns = ['f1_1','f1_2','f2_1','f2_2', 'op'])
test_df = pd.DataFrame(test, columns = ['f1_1','f1_2','f2_1','f2_2', 'op'])
tc_Xtrain = train_df[train_df.columns[~train_df.columns.isin(['op'])]]
tc_ytrain = train_df['op'].astype(np.float32)
tc_Xtest = test_df[test_df.columns[~test_df.columns.isin(['op'])]]
tc_ytest = test_df['op'].astype(np.float32)

#normalize values
scaler = StandardScaler()
tc_Xtrain = scaler.fit_transform(tc_Xtrain)
# print(tc_Xtrain)
# print(f'Type after standard scaler = {type(tc_Xtrain)}')
tc_Xtest = scaler.transform(tc_Xtest)
X_train = np.apply_along_axis(create_ts_features, 1, tc_Xtrain, num_features=2).astype(np.float32)
X_test = np.apply_along_axis(create_ts_features, 1, tc_Xtest, num_features=2).astype(np.float32)
# print(X_train)
# print(tc_ytrain)

train_lstm(X_train, tc_ytrain, X_test, tc_ytest, num_epochs=4000, f1_thresh=0.6, batch_size=100)

指标计算代码

def plot_pr_withtorch(target, pred, fig, title):
    pr_curve = PrecisionRecallCurve(pos_label=1)
    precision, recall, thresholds = pr_curve(pred, target)
    print_key_prs(precision, recall, thresholds, title)
    N = len(recall)
    fig.add_trace(go.Scatter(x=recall[0 : N - 1], y=precision[0 : N - 1], name=title))


def get_scores(y, y_preds, print_label):
    print(f"{print_label} summary:")
    confusion = confusion_matrix(y, y_preds)
    print(f"Confusion matrix: {confusion}")
    print("Accuracy: {:.2f}".format(accuracy_score(y, y_preds)))
    print("Precision: {:.2f}".format(precision_score(y, y_preds)))
    print("Recall: {:.2f}".format(recall_score(y, y_preds)))
    F1_score = f1_score(y, y_preds)
    print("F1: {:.2f}".format(f1_score(y, y_preds)))
    return F1_score

技术问题解答

1. 不同Batch Size下模型表现差异的直观原理

  • 梯度噪声的引导作用:Batch Size=1时,每次更新用的是单样本的梯度,存在较大噪声,但这种噪声能帮助模型跳出局部最优,找到更泛化的参数。而全量Batch的梯度是所有样本的平均,方向更“平滑”,在小数据集(仅12个样本)中容易陷入无法区分两类的局部最优。
  • 参数更新频率差异:Batch Size=1时,每个epoch内参数更新12次,而全量Batch每个epoch仅更新1次。高频更新让模型能更快对每个样本的模式做出响应,快速调整参数以拟合0类和1类的区分特征。
  • 小数据集的信号稀释:当训练样本极少时,全量Batch的平均梯度会稀释单个样本的有效分类信号,导致模型无法学到清晰的边界;而单样本更新能让模型逐个聚焦样本特征,精准调整参数。

2. 训练数据量与Batch Size选择的经验法则,以及小Batch的收敛风险

经验法则

  • 小数据集(<1000样本):优先选小Batch(1-32)甚至单样本,利用梯度噪声避免过拟合或局部最优,同时高频更新能快速学习有限样本中的模式。
  • 中等数据集(1000-10000样本):常用Batch Size为32-128,平衡更新频率和计算效率。
  • 大数据集(>10000样本):可选择更大的Batch Size(256-1024),此时梯度平均噪声足够小,大Batch能提升计算效率且不显著影响性能。
  • 通用参考:Batch Size通常选2的幂次(如16、32、64),方便GPU内存调度,但并非硬性要求。

小Batch的收敛风险

小Batch确实会增加训练时间,但不会必然提升不收敛概率:

  • 收敛性:只要学习率设置合理,小Batch训练最终能收敛,甚至可能收敛到更优的局部最优。若学习率过大,梯度噪声可能导致参数震荡,此时需适当降低学习率或使用学习率调度器(如ReduceLROnPlateau)。
  • 训练稳定性:小Batch的损失曲线波动更大,但这是正常的噪声波动,并非不收敛。可通过观察多个epoch的平均损失判断是否真正收敛。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 03:35:22