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

如何在TensorFlow中使用GPU进行深度学习模型训练?

问题背景

计算机环境

系统与CUDA信息

Microsoft Windows [Version 10.0.22621.963]
(c) Microsoft Corporation. All rights reserved.

C:\Users\donhu>nvcc -V
nvcc: NVIDIA (R) Cuda compiler driver
Copyright (c) 2005-2022 NVIDIA Corporation
Built on Tue_May__3_19:00:59_Pacific_Daylight_Time_2022
Cuda compilation tools, release 11.7, V11.7.64
Build cuda_11.7.r11.7/compiler.31294372_0

C:\Users\donhu>nvidia-smi
Sat Dec 17 23:40:44 2022
+-----------------------------------------------------------------------------+
| NVIDIA-SMI 512.77       Driver Version: 512.77       CUDA Version: 11.6     |
|-------------------------------+----------------------+----------------------+| GPU  Name            TCC/WDDM | Bus-Id        Disp.A | Volatile Uncorr. ECC |
| Fan  Temp  Perf  Pwr:Usage/Cap|         Memory-Usage | GPU-Util  Compute M. |
|                               |                      |               MIG M. |
|===============================+======================+======================|
|   0  NVIDIA GeForce ... WDDM  | 00000000:01:00.0  On |                  N/A |
| 34%   31C    P8    16W / 125W |   1377MiB /  6144MiB |      4%      Default |
|                               |                      |                  N/A |
+-------------------------------+----------------------+----------------------+

+-----------------------------------------------------------------------------+
| Processes:                                                                  |
|  GPU   GI   CI        PID   Type   Process name                  GPU Memory |
|        ID   ID                                                   Usage      |
|=============================================================================|
|    0   N/A  N/A      3392    C+G   C:\Windows\explorer.exe         N/A      |
|    0   N/A  N/A      4484    C+G   ...artMenuExperienceHost.exe    N/A      |
|    0   N/A  N/A      6424    C+G   ...n1h2txyewy\SearchHost.exe    N/A      |
|    0   N/A  N/A      6796    C+G   ...lPanel\SystemSettings.exe    N/A      |
|    0   N/A  N/A      7612    C+G   ...8bbwe\WindowsTerminal.exe    N/A      |
|    0   N/A  N/A      9700    C+G   ...8bbwe\WindowsTerminal.exe    N/A      |
|    0   N/A  N/A     10624    C+G   ...perience\NVIDIA Share.exe    N/A      |
|    0   N/A  N/A     10728    C+G   ...er Java\jre\bin\javaw.exe    N/A      |
|    0   N/A  N/A     13064    C+G   ...8bbwe\WindowsTerminal.exe    N/A      |
|    0   N/A  N/A     14496    C+G   ...462.46\msedgewebview2.exe    N/A      |
|    0   N/A  N/A     17124    C+G   ...ooting 2\BugShooting2.exe    N/A      |
|    0   N/A  N/A     19064    C+G   ...8bbwe\Notepad\Notepad.exe    N/A      |
|    0   N/A  N/A     19352    C+G   ...8bbwe\WindowsTerminal.exe    N/A      |
|    0   N/A  N/A     20920    C+G   ...y\ShellExperienceHost.exe    N/A      |
|    0   N/A  N/A     21320    C+G   ...e\PhoneExperienceHost.exe    N/A      |
|    0   N/A  N/A     21368    C+G   ...me\Application\chrome.exe    N/A      |
+-----------------------------------------------------------------------------+

C:\Users\donhu>

训练代码

from tensorflow import keras
from tensorflow.keras import layers

def get_model():
    model = keras.Sequential([
        layers.Dense(512, activation="relu"),
        layers.Dense(10, activation="softmax")
    ])
    model.compile(optimizer="rmsprop",
                  loss="sparse_categorical_crossentropy",
                  metrics=["accuracy"])
    return model

model = get_model()
history_noise = model.fit(
    train_images_with_noise_channels, train_labels,
    epochs=10,
    batch_size=128,
    validation_split=0.2)

model = get_model()
history_zeros = model.fit(
    train_images_with_zeros_channels, train_labels,
    epochs=10,
    batch_size=128,
    validation_split=0.2)

问题

请问如何在TensorFlow中使用GPU进行模型训练?


解决方案

1. 匹配TensorFlow与CUDA版本

你的环境是CUDA 11.7,需安装适配的TensorFlow版本,推荐安装2.10.x或2.11.x版本,执行以下命令完成安装:

pip install tensorflow==2.10.0

同时确保已安装对应版本的cuDNN(CUDA 11.7适配cuDNN 8.4.x),并配置好系统环境变量:将CUDA的bin和libnvvp目录添加到系统PATH,设置CUDA_PATH指向CUDA安装根目录。

2. 验证GPU是否被TensorFlow识别

在代码开头加入以下片段,确认GPU可用性:

import tensorflow as tf
print("GPU可用状态:", tf.test.is_gpu_available())
print("检测到的GPU设备:", tf.config.list_physical_devices('GPU'))

若输出显示GPU可用,说明环境配置正确;若未检测到,检查CUDA、cuDNN的安装路径及环境变量是否配置无误。

3. 启用GPU训练

  • 自动启用:TensorFlow默认优先使用GPU,你现有的训练代码无需修改,model.fit()会自动在GPU上执行计算。
  • 手动指定GPU:如果有多块GPU,可指定使用某一块:
gpus = tf.config.list_physical_devices('GPU')
if gpus:
    tf.config.set_visible_devices(gpus[0], 'GPU')  # 指定使用第0块GPU
  • 限制GPU内存:避免TensorFlow占满GPU内存,可设置内存按需增长:
gpus = tf.config.list_physical_devices('GPU')
if gpus:
    try:
        for gpu in gpus:
            tf.config.experimental.set_memory_growth(gpu, True)
    except RuntimeError as e:
        print(e)

4. 确认训练正在使用GPU

  • 打开Windows任务管理器,切换到“性能”标签,查看GPU使用率,训练过程中使用率明显上升即说明正在使用GPU。
  • 在代码中打印模型参数所在设备,验证是否在GPU上:
model = get_model()
print("模型参数所在设备:", model.layers[0].weights[0].device)

内容的提问来源于stack exchange,提问作者Đỗ Như Vỹ

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 11:45:33