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

如何确保timeseriesAI/tsai使用GPU进行时序分类训练

如何确保tsai 0.3.5使用GPU进行时序分类训练

1. 先确认GPU是否被PyTorch识别

tsai基于PyTorch运行,首先要验证环境中PyTorch能否调用GPU。在代码开头添加以下检查:

import torch
# 检查GPU可用性
print("GPU可用:", torch.cuda.is_available())
if torch.cuda.is_available():
    device = torch.device("cuda:0")
    print("当前使用GPU:", torch.cuda.get_device_name(0))
else:
    device = torch.device("cpu")
    print("当前使用CPU")

如果输出GPU可用: False,说明你的PyTorch是CPU版本,需要卸载后安装对应CUDA版本的PyTorch。

2. 强制tsai使用GPU训练

tsai的TSClassifier默认会自动检测GPU,但如果未生效,可手动指定设备,同时确保数据格式适配GPU:

修改后的完整代码

import os
os.chdir(os.path.dirname(os.path.abspath(__file__)))
from pickle import load
import numpy as np
import torch
from tsai.all import *
import matplotlib.pyplot as plt
from sklearn.metrics import precision_recall_curve

# 检查并设置设备
print("GPU可用:", torch.cuda.is_available())
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
print(f"使用设备: {device}")

num_datesets = 1
for dataset_idx in range(num_datesets):
    X_train = load(open(r"X_train_"+str(dataset_idx)+".pkl", 'rb'))
    y_train = load(open(r"y_train_"+str(dataset_idx)+".pkl", 'rb'))
    X_test = load(open(r"X_test_"+str(dataset_idx)+".pkl", 'rb'))
    y_test = load(open(r"y_test_"+str(dataset_idx)+".pkl", 'rb'))
    print("dataset loaded")

    # 转换数据为float32(GPU更适配该格式)
    X_train = X_train.astype(np.float32)
    X_test = X_test.astype(np.float32)

    # 创建分类器时指定设备
    learn = TSClassifier(X_train, y_train, arch=InceptionTimePlus, 
                         arch_config=dict(fc_dropout=0.5), device=device)

    print("training started")
    learn.fit_one_cycle(5, 0.0005)
    learn.export("tsai_"+str(dataset_idx)+".pkl") 
    
    probas, target, preds = learn.get_X_preds(X_test, y_test)
    precision, recall, thresholds = precision_recall_curve(target, probas)
    plt.clf()
    plt.fill_between(recall, precision)
    plt.ylabel("Precision")
    plt.xlabel("Recall")
    plt.title("tsai_"+str(dataset_idx)+"_precision_recall_curve")
    plt.savefig("tsai_"+str(dataset_idx)+".png")
    plt.show()

3. 额外检查项

  • 确保PyTorch版本与你的CUDA版本匹配,可通过torch.version.cuda查看PyTorch绑定的CUDA版本
  • 代码中未使用的multiprocessing.Process可移除,避免潜在冲突
  • 训练时用nvidia-smi命令观察GPU使用率,若使用率低,可能是数据加载瓶颈,可优化数据加载流程

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 07:57:36