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

Kaggle中TPU/GPU未生效,LSTM自编码器仅用CPU训练求助

问题分析与解决方案

核心错误:模型定义未在TPU策略作用域内

你的代码中with tpu_strategy.scope():存在缩进错误,仅变量n_steps和n_features处于TPU策略作用域内,模型创建、编译的核心代码都在作用域之外,导致模型默认在CPU上初始化,无法调用TPU训练。

修正后的完整代码

import numpy as np
import tensorflow as tf
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Input, LSTM, Dense, RepeatVector, TimeDistributed
from sklearn.model_selection import train_test_split

# 加载数据
data = np.load("/kaggle/working/dataStack.npy")
# 从数据维度获取时间步长,假设data形状为(样本数, 时间步, 特征数)
maxlen = data.shape[1]

# TPU初始化与策略设置
tpu = tf.distribute.cluster_resolver.TPUClusterResolver()
tf.tpu.experimental.initialize_tpu_system(tpu)
tpu_strategy = tf.distribute.TPUStrategy(tpu)

# 所有模型相关操作必须放在TPU策略作用域内
with tpu_strategy.scope():
    n_steps = maxlen
    n_features = 3  # position, torque, thrust

    # 定义LSTM自编码器
    model = Sequential()
    model.add(Input(shape=(n_steps, n_features)))
    model.add(LSTM(128, activation='relu'))
    model.add(Dense(64, activation='relu', kernel_initializer='he_uniform'))
    model.add(RepeatVector(n_steps))
    model.add(LSTM(128, activation='relu', return_sequences=True))
    model.add(TimeDistributed(Dense(n_features)))

    model.compile(optimizer='adam', loss='mse', steps_per_execution=32)

# 自编码器输入输出一致,X和y均为原始数据
X = data
y = data
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

# 计算适配TPU的批量大小
BATCH_SIZE = 16 * tpu_strategy.num_replicas_in_sync

# 启动训练
model.fit(X_train, y_train, epochs=10, batch_size=BATCH_SIZE, validation_data=(X_test, y_test))

额外注意事项

  1. 变量完整性:原代码中X未赋值、maxlen未定义,修正后从数据维度自动获取时间步长,保证变量合法
  2. 日志警告处理:你看到的SetPriority unimplemented是TensorFlow与TPU交互的正常日志,不影响训练,可忽略
  3. TPU验证:训练启动后,查看Kaggle右侧"Accelerator"面板,会显示TPU占用状态;也可通过model.summary()查看模型的设备分配信息

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 09:32:34