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

将TensorFlow 1线性回归代码迁移至TensorFlow 2(无sess.run实现)

从TensorFlow 1迁移到TensorFlow 2:简洁线性回归实现

你在迁移TensorFlow 1的线性回归代码到TF2时遇到的核心问题,是没搞清楚TF2的**即刻执行模式(Eager Execution)**和TF1静态图模式的差异——TF2不需要占位符(placeholder)和会话(session),而是直接用Python动态执行逻辑,配合tf.GradientTape来追踪梯度。

你的TF2尝试里把输入X和y定义成tf.Variable是错误的,这两个是训练数据,不是需要优化的参数,完全没必要做成可训练变量。下面是修正后的简洁实现,完美对应你原来TF1代码的逻辑:

%tensorflow_version 2.x
import numpy as np
import tensorflow as tf

# 保持原有的学习常量设置
n_samples, batch_size, num_steps = 1000, 100, 20000
display_step = 100

# 生成训练数据,和TF1版本完全一致
X_data = np.random.uniform(1, 10, (n_samples, 1))
y_data = 2 * X_data + 1 + np.random.normal(0, 2, (n_samples, 1))

# 定义需要优化的参数k和b——这两个才是可训练变量
k = tf.Variable(tf.random.normal((1, 1)), name='slope')
b = tf.Variable(tf.zeros((1,)), name='bias')

# 选择SGD优化器,学习率和TF1版本一致(0.0001,你之前试的0.01太大易震荡)
optimizer = tf.keras.optimizers.SGD(learning_rate=0.0001)

# 训练循环
for i in range(num_steps):
    # 随机选取批量样本,逻辑和TF1一致
    indices = np.random.choice(n_samples, batch_size)
    X_batch, y_batch = X_data[indices], y_data[indices]
    
    # 用GradientTape追踪梯度:包裹前向计算和损失计算
    with tf.GradientTape() as tape:
        # 前向传播计算预测值
        y_pred = tf.matmul(X_batch, k) + b
        # 计算损失(均方误差的总和,和TF1一致)
        loss_val = tf.reduce_sum((y_batch - y_pred) ** 2)
    
    # 获取损失对k和b的梯度
    grads = tape.gradient(loss_val, [k, b])
    
    # 用优化器更新参数:把梯度和对应的变量配对传入
    optimizer.apply_gradients(zip(grads, [k, b]))
    
    # 定期打印训练状态
    if (i + 1) % display_step == 0:
        print(f'Epoch {i+1}: loss={loss_val.numpy():.8f}, k={k.numpy()[0][0]:.4f}, b={b.numpy()[0]:.4f}')

关键要点解释:

  • 移除占位符和会话:TF2的即刻执行模式下,不需要tf.placeholder和tf.Session,直接用numpy数组作为输入传入计算即可。
  • 正确使用tf.GradientTape:GradientTape会记录上下文内的所有张量运算,从而自动计算损失对可训练变量(k和b)的梯度。
  • 手动更新参数:通过optimizer.apply_gradients方法,把梯度和对应的变量配对传入,完成参数更新——这对应TF1里sess.run(optimizer, feed_dict)的逻辑。
  • 保持批量采样逻辑:和你原来的TF1代码完全一致,每次迭代随机选择batch样本,实现随机梯度下降。

运行这段代码后,你会看到k逐渐逼近2,b逐渐逼近1,和TF1版本的效果完全一致,同时代码更加简洁直观。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.09 17:18:15