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

Python中如何用大数据集训练Random Forest分类器避免内存错误

大内存数据集训练Random Forest内存耗尽问题优化方案

问题背景

我有一个3000万行的数据集,包含两列:一列是0/1标签,另一列每行有1280个特征,总大小181GB。用Random Forest训练时内存耗尽崩溃(400GB内存也不行)。数据集是Hugging Face arrow格式,加载后转成DataFrame处理,怀疑这一步占了太多内存。知道可以降维,但想先通过修改代码降低内存占用,原代码如下:

import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
from sklearn.model_selection import train_test_split
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import accuracy_score, roc_auc_score, roc_curve, auc
from datasets import load_dataset, Dataset

# Load dataset
df = Dataset.from_file("data.arrow")
df = pd.DataFrame(df)
X = df['embeddings'].to_numpy() # Convert Series to NumPy array
X = np.array(X.tolist()) # Convert list of arrays to a 2D NumPy array
X = X.reshape(X.shape[0], -1) # Flatten the 3D array into a 2D array
y = df['labels']

# Split the data into training and testing sets (80% train, 20% test)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

# Initialize the random forest classifier
rf_classifier = RandomForestClassifier(n_estimators=100, random_state=42)

# Train the classifier
rf_classifier.fit(X_train, y_train)

# Make predictions on the test set
y_pred = rf_classifier.predict(X_test)

# Evaluate the classifier

# Calculate accuracy
accuracy = accuracy_score(y_test, y_pred)
print("Accuracy:", accuracy)

# Calculate AUC score
auc_score = roc_auc_score(y_test, y_pred)
print("AUC Score:", auc_score)

with open("metrics.txt", "w") as f:
    f.write("Accuracy: " + str(accuracy) + "\n")
    f.write("AUC Score: " + str(auc_score))
    
# Make predictions on the test set
y_pred_proba = rf_classifier.predict_proba(X_test)[:, 1]

# Calculate ROC curve
fpr, tpr, thresholds = roc_curve(y_test, y_pred_proba)

# Plot ROC curve
plt.figure()
plt.plot(fpr, tpr, color='darkorange', lw=2, label='ROC curve (area = %0.2f)' % auc(fpr, tpr))
plt.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--')
plt.xlim([0.0, 1.0])
plt.ylim([0.0, 1.05])
plt.xlabel('False Positive Rate')
plt.ylabel('True Positive Rate')
plt.title('Receiver Operating Characteristic (ROC) Curve')
plt.legend(loc="lower right")

# Save ROC curve plot to an image file
plt.savefig('roc_curve.png')

# Close plot to free memory

核心优化措施

1. 彻底放弃DataFrame转换

转成DataFrame会强制把整个数据集加载到内存,这是内存爆炸的核心原因。直接用Hugging Face Dataset的原生接口处理,它支持按需加载和分批处理,不用一次性读入所有数据:

# 直接加载arrow数据集,不转DataFrame
dataset = load_dataset("arrow", data_files="data.arrow", split="train")
# 直接在Dataset层面拆分训练测试集,不用全量加载后拆分
dataset_split = dataset.train_test_split(test_size=0.2, random_state=42)
train_dataset = dataset_split['train']
test_dataset = dataset_split['test']

2. 优化特征转换,减少冗余复制

原代码里的特征转换做了三次冗余操作(转numpy→转列表→转numpy→reshape),每次都会生成新数组,额外占用内存。用map批量处理特征,直接生成低精度的2D数组:

def process_batch(batch):
    # 直接堆叠embeddings成2D数组,同时转成float32(内存减半)
    batch['embeddings'] = np.stack(batch['embeddings']).reshape(len(batch), -1).astype(np.float32)
    return batch

# 批量处理,避免一次性加载所有特征
train_dataset = train_dataset.map(process_batch, batched=True, batch_size=2000)
test_dataset = test_dataset.map(process_batch, batched=True, batch_size=2000)

3. 用分布式/分批训练的Random Forest替代原生版本

Scikit-learn的原生RandomForest不支持分批训练,必须全量加载数据。可以用dask-ml的RandomForest,它支持大数据集的分批处理,不需要把所有数据塞进内存:

from dask_ml.ensemble import RandomForestClassifier

# 转成Dask DataFrame,自动支持分批加载
train_dask = train_dataset.to_dask()
X_train = train_dask['embeddings']
y_train = train_dask['labels']

# 初始化Dask版RF,训练时自动分批处理数据
rf_classifier = RandomForestClassifier(n_estimators=100, random_state=42)
rf_classifier.fit(X_train, y_train)

4. 按需加载测试集数据

测试集的标签和预测结果可以按需加载,不用一次性读入所有数据:

test_dask = test_dataset.to_dask()
X_test = test_dask['embeddings']
# 标签数据量小,直接加载到内存即可
y_test = test_dask['labels'].compute()

# 预测结果按需计算
y_pred = rf_classifier.predict(X_test).compute()
y_pred_proba = rf_classifier.predict_proba(X_test)[:, 1].compute()

完整优化后代码

import numpy as np
import matplotlib.pyplot as plt
from sklearn.metrics import accuracy_score, roc_auc_score, roc_curve, auc
from datasets import load_dataset
from dask_ml.ensemble import RandomForestClassifier

# 加载数据集,不转DataFrame
dataset = load_dataset("arrow", data_files="data.arrow", split="train")
# 拆分训练测试集
dataset_split = dataset.train_test_split(test_size=0.2, random_state=42)
train_dataset = dataset_split['train']
test_dataset = dataset_split['test']

# 批量处理特征,转换为float32的2D数组
def process_batch(batch):
    batch['embeddings'] = np.stack(batch['embeddings']).reshape(len(batch), -1).astype(np.float32)
    return batch

train_dataset = train_dataset.map(process_batch, batched=True, batch_size=2000)
test_dataset = test_dataset.map(process_batch, batched=True, batch_size=2000)

# 转换为Dask DataFrame,支持分批训练
train_dask = train_dataset.to_dask()
X_train = train_dask['embeddings']
y_train = train_dask['labels']

# 初始化Dask版Random Forest
rf_classifier = RandomForestClassifier(n_estimators=100, random_state=42)
# 训练(自动分批处理,不占满内存)
rf_classifier.fit(X_train, y_train)

# 处理测试集
test_dask = test_dataset.to_dask()
X_test = test_dask['embeddings']
y_test = test_dask['labels'].compute() # 标签量小,直接加载到内存

# 预测
y_pred = rf_classifier.predict(X_test).compute()
y_pred_proba = rf_classifier.predict_proba(X_test)[:, 1].compute()

# 评估指标
accuracy = accuracy_score(y_test, y_pred)
auc_score = roc_auc_score(y_test, y_pred)
print(f"Accuracy: {accuracy}")
print(f"AUC Score: {auc_score}")

with open("metrics.txt", "w") as f:
    f.write(f"Accuracy: {accuracy}\n")
    f.write(f"AUC Score: {auc_score}")

# 绘制ROC曲线
fpr, tpr, thresholds = roc_curve(y_test, y_pred_proba)
plt.figure()
plt.plot(fpr, tpr, color='darkorange', lw=2, label=f'ROC curve (area = {auc(fpr, tpr):.2f})')
plt.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--')
plt.xlim([0.0, 1.0])
plt.ylim([0.0, 1.05])
plt.xlabel('False Positive Rate')
plt.ylabel('True Positive Rate')
plt.title('Receiver Operating Characteristic (ROC) Curve')
plt.legend(loc="lower right")
plt.savefig('roc_curve.png')
plt.close()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 18:38:11