这次我们不聊大模型,也不聊复杂的深度学习。直接回到机器学习最基础、也是 AI 开发中最常用到的模型:线性回归。

如果你正在入门 AI 开发,或者想把机器学习模型集成到 Web 项目里,线性回归是最适合作为起点的算法。它原理直观、代码实现难度低、对硬件几乎没有要求,又能完整走一遍“数据准备 -> 模型训练 -> 模型评估 -> 接口部署 -> 前端调用”的全流程。

这篇文章会从原理、Python 手写实现、sklearn 快速实现,一路做到 Flask 接口封装和简单前端页面调用。全程基于 Web 技术栈,不需要独立显卡,不需要昂贵服务器,普通开发机就能跑完。最后还会给出常见报错排查和工程化建议。

1. 核心能力速览

能力项 说明
项目类型 机器学习入门算法 + Web 接口部署演示
核心算法 线性回归(Linear Regression)
依赖环境 Python 3.8+,Flask、scikit-learn、numpy、pandas
硬件要求 CPU 即可,无显卡要求
训练数据 示例数据或自定义 CSV 数据
Web 能力 Flask 提供 REST API,前端 HTML 页面调用
批量任务 支持批量预测,一次请求可传入多条数据
适合人群 Web 开发者、AI 开发初学者、需要构建预测接口的工程师
项目形态 教学演示项目,可扩展到真实业务预测场景

这个项目最大的价值不是算法本身,而是把线性回归从“数学公式”变成“可运行的 Web 服务”。你会看到一条完整的链路:模型文件如何保存、如何加载、如何通过 HTTP 接口对外提供服务、前端如何拿到预测结果。

2. 适用场景与使用边界

线性回归解决的核心问题是: 根据已知数据推断变量之间的线性关系,并用这个关系做预测

典型场景包括:

  • 房价预测:根据面积、楼层、位置等特征预测价格。
  • 销量预估:根据历史销量、促销力度、季节等因素预估未来销量。
  • 广告投放效果分析:评估不同渠道投入与转化率的关系。
  • Web 应用中的趋势分析与数据监控。

但线性回归有明确的使用边界:

  • 它假设特征与目标值之间是线性关系。如果数据呈现明显的曲线或周期性,线性回归效果会很差。
  • 对异常值敏感。极端值会显著拉高损失,需要做数据清洗。
  • 不擅长处理高维稀疏数据。文本分类、图像识别这类任务不要用线性回归。
  • 若特征之间存在高度多重共线性,模型参数会不稳定。

另外需要提醒:如果训练数据涉及个人信息、用户行为、商业数据,必须在获得授权的前提下使用。模型部署到 Web 服务时,不要直接暴露在公网无防护环境下,至少要做接口访问控制或部署在内网。

3. 环境准备与前置条件

本项目不需要 GPU,不需要 CUDA,也不需要 Docker。你只需要一个能运行 Python 的开发环境。

3.1 确认 Python 版本

在终端执行:

python --version

推荐使用 Python 3.8 到 3.11 版本。如果版本过低,建议先升级。

3.2 创建独立虚拟环境

为了避免和系统 Python 环境冲突,建议创建虚拟环境。

# 创建虚拟环境
python -m venv ml_env

# 激活虚拟环境
# Windows
ml_env\Scripts\activate

# macOS / Linux
source ml_env/bin/activate

激活成功的标志是命令行前缀出现 (ml_env)

3.3 安装依赖库

pip install numpy pandas scikit-learn flask

安装完成后,验证关键库是否可用:

python -c "import sklearn; print(sklearn.__version__)"
python -c "import flask; print(flask.__version__)"

如果输出版本号,说明环境就绪。

3.4 准备目录结构

建议按下面的结构组织项目文件:

linear_regression_web/
├── app.py                 # Flask Web 服务
├── train_model.py         # 模型训练脚本
├── models/
│   └── linear_model.pkl   # 训练好的模型文件
├── static/
│   └── index.html         # 前端页面
└── data/
    └── house_data.csv     # 训练数据

这种目录结构的好处是:训练代码、服务代码、模型文件、前端资源分开放,后续扩展和维护都比较清晰。

4. 线性回归原理与数学基础

线性回归的数学形式非常简单:

y = w1 * x1 + w2 * x2 + ... + wn * xn + b

其中:

  • x1, x2, ..., xn 是特征。
  • w1, w2, ..., wn 是每个特征对应的权重。
  • b 是偏置项。
  • y 是预测目标。

训练过程的本质是:找到一组 w b ,使得预测值 y_pred 与真实值 y_true 之间的差距最小。

这个“差距”通常用**均方误差(MSE)**来衡量:

MSE = (1/n) * sum((y_true - y_pred)^2)

MSE 越小,说明模型预测越准。

为了让 MSE 尽可能小,最常用的优化方法是 梯度下降 。简单理解:计算当前参数下的梯度方向,沿着梯度下降的方向更新参数,不断迭代直到损失收敛。

梯度下降中有一个关键参数叫 学习率(learning rate) ,通常用 0.01 0.001 这样的值。学习率太大,参数会震荡甚至发散;学习率太小,收敛速度很慢。

5. 用 Python 手动实现线性回归

在引入 sklearn 之前,先用 numpy 手写一个简单的线性回归。这一步的目的是理解内部计算逻辑,而不是直接黑盒调库。

5.1 手写梯度下降版本

import numpy as np

# 模拟数据:y = 2 * x + 1 + 噪声
np.random.seed(42)
X = np.random.rand(100, 1) * 10
y = 2 * X + 1 + np.random.randn(100, 1) * 1.5

# 初始化参数
w = np.random.randn(1, 1)
b = np.random.randn(1)

learning_rate = 0.01
epochs = 1000
n = len(X)

for epoch in range(epochs):
    y_pred = X @ w + b
    loss = np.mean((y - y_pred) ** 2)

    # 计算梯度
    dw = (-2 / n) * (X.T @ (y - y_pred))
    db = (-2 / n) * np.sum(y - y_pred)

    # 更新参数
    w -= learning_rate * dw
    b -= learning_rate * db

    if epoch % 200 == 0:
        print(f"Epoch {epoch}, Loss: {loss:.4f}")

print(f"最终权重 w = {w.flatten()[0]:.4f}")
print(f"最终偏置 b = {b.flatten()[0]:.4f}")

运行这段代码,你会看到损失逐渐下降,最终权重接近 2 ,偏置接近 1 。这就是线性回归训练的完整逻辑。

5.2 为什么需要 sklearn

手写版本适合理解原理,但工程项目中直接用 sklearn 更高效。sklearn 的 LinearRegression 内部使用最小二乘法求解,不需要手动调整学习率,训练速度更快,结果也更稳定。

6. 使用 scikit-learn 实现线性回归

下面的代码会完成:数据加载、数据集划分、模型训练、模型保存四个步骤。

6.1 生成示例数据

import numpy as np
import pandas as pd

np.random.seed(42)
n = 200

# 模拟房价数据:面积、房间数、楼龄 -> 价格
area = np.random.randint(40, 150, n)
rooms = np.random.randint(1, 5, n)
age = np.random.randint(0, 30, n)

price = 0.8 * area + 5 * rooms - 1.2 * age + 30 + np.random.randn(n) * 8

df = pd.DataFrame({
    "area": area,
    "rooms": rooms,
    "age": age,
    "price": price
})

df.to_csv("data/house_data.csv", index=False)
print(df.head())

6.2 训练并保存模型

import pandas as pd
from sklearn.model_selection import train_test_split
from sklearn.linear_model import LinearRegression
from sklearn.metrics import mean_squared_error, r2_score
import joblib

# 加载数据
data = pd.read_csv("data/house_data.csv")

# 特征与目标
X = data[["area", "rooms", "age"]]
y = data["price"]

# 划分训练集和测试集
X_train, X_test, y_train, y_test = train_test_split(
    X, y, test_size=0.2, random_state=42
)

# 创建并训练模型
model = LinearRegression()
model.fit(X_train, y_train)

# 预测
y_pred = model.predict(X_test)

# 评估
mse = mean_squared_error(y_test, y_pred)
r2 = r2_score(y_test, y_pred)

print(f"线性回归模型训练完成")
print(f"MSE: {mse:.4f}")
print(f"R2 Score: {r2:.4f}")
print(f"模型系数: {model.coef_}")
print(f"模型截距: {model.intercept_:.4f}")

# 保存模型
joblib.dump(model, "models/linear_model.pkl")
print("模型已保存到 models/linear_model.pkl")

R2 Score 是回归模型的常用评估指标,取值范围通常为 0 到 1。越接近 1,说明模型对数据的拟合程度越好。在这个模拟数据上,R2 一般会在 0.9 以上。

到此,模型训练部分完成。接下来是把模型变成 Web 服务。

7. Flask 接口 API 部署

Flask 是 Python 生态中最轻量的 Web 框架之一,适合快速把模型封装成 HTTP 接口。

7.1 创建 Flask 应用

在项目根目录创建 app.py

import joblib
import numpy as np
import pandas as pd
from flask import Flask, request, jsonify

app = Flask(__name__)

# 加载模型
model = joblib.load("models/linear_model.pkl")

@app.route("/health", methods=["GET"])
def health():
    return jsonify({"status": "ok", "message": "linear regression service is running"})

@app.route("/predict", methods=["POST"])
def predict():
    try:
        # 获取请求数据
        data = request.get_json()

        # 兼容单条和批量请求
        if "features" in data:
            # 批量预测模式
            features_list = data["features"]
            df = pd.DataFrame(features_list)
        else:
            # 单条预测模式
            df = pd.DataFrame([data])

        required_cols = ["area", "rooms", "age"]
        for col in required_cols:
            if col not in df.columns:
                return jsonify({"error": f"缺少必要特征: {col}"}), 400

        # 执行预测
        predictions = model.predict(df[required_cols])

        return jsonify({
            "code": 0,
            "predictions": predictions.tolist()
        })

    except Exception as e:
        return jsonify({"error": str(e)}), 500

if __name__ == "__main__":
    app.run(host="127.0.0.1", port=5000, debug=False)

7.2 启动服务

python app.py

启动成功后,控制台会显示 Running on http://127.0.0.1:5000

确认服务健康:

curl http://127.0.0.1:5000/health

预期返回:

{"message":"linear regression service is running","status":"ok"}

7.3 接口调用示例

单条预测
curl -X POST http://127.0.0.1:5000/predict \
  -H "Content-Type: application/json" \
  -d "{\"area\": 85, \"rooms\": 2, \"age\": 5}"

预期返回:

{
  "code": 0,
  "predictions": [98.871]
}
批量预测
curl -X POST http://127.0.0.1:5000/predict \
  -H "Content-Type: application/json" \
  -d "{\"features\": [{\"area\": 85, \"rooms\": 2, \"age\": 5}, {\"area\": 120, \"rooms\": 4, \"age\": 3}]}"

预期返回:

{
  "code": 0,
  "predictions": [98.871, 139.495]
}

这里 predictions 数组里的每个元素对应该请求中一条数据的预测结果。

7.4 Python 客户端调用

import requests

url = "http://127.0.0.1:5000/predict"

# 单条预测
payload = {
    "area": 100,
    "rooms": 3,
    "age": 10
}

response = requests.post(url, json=payload, timeout=10)
print(response.json())

接口能跑通,后面就可以接到自己的业务系统里。无论是小程序后端、数据中台,还是自动化报表工具,都可以把预测结果通过这个接口做集成。

8. 前端页面调用接口

为了让整个流程更完整,写一个简单的 HTML 页面,通过浏览器调用 Flask 接口完成预测。

创建 static/index.html

<!DOCTYPE html>
<html lang="zh-CN">
<head>
    <meta charset="UTF-8">
    <meta name="viewport" content="width=device-width, initial-scale=1.0">
    <title>线性回归房价预测</title>
    <style>
        body { font-family: "Microsoft YaHei", sans-serif; max-width: 600px; margin: 40px auto; padding: 20px; background: #f7f8fa; }
        h1 { text-align: center; }
        .form-group { margin-bottom: 16px; }
        label { display: block; margin-bottom: 6px; font-weight: bold; }
        input[type="number"] { width: 100%; padding: 10px; border: 1px solid #ccc; border-radius: 4px; font-size: 16px; }
        button { width: 100%; padding: 12px; background: #1a73e8; color: #fff; border: none; border-radius: 4px; font-size: 16px; cursor: pointer; }
        button:hover { background: #1558b0; }
        .result { margin-top: 20px; padding: 20px; background: #fff; border-radius: 4px; text-align: center; font-size: 20px; }
        .error { color: #d93025; margin-top: 12px; text-align: center; }
    </style>
</head>
<body>
    <h1>房价预测</h1>
    <form id="prediction-form">
        <div class="form-group">
            <label for="area">面积(平方米)</label>
            <input type="number" id="area" value="86" required>
        </div>
        <div class="form-group">
            <label for="rooms">房间数</label>
            <input type="number" id="rooms" value="2" required>
        </div>
        <div class="form-group">
            <label for="age">楼龄(年)</label>
            <input type="number" id="age" value="5" required>
        </div>
        <button type="submit">开始预测</button>
    </form>
    <div class="result" id="result">等待输入</div>
    <div class="error" id="error"></div>

    <script>
        document.getElementById("prediction-form").addEventListener("submit", async (e) => {
            e.preventDefault();

            const payload = {
                area: parseFloat(document.getElementById("area").value),
                rooms: parseFloat(document.getElementById("rooms").value),
                age: parseFloat(document.getElementById("age").value)
            };

            try {
                const response = await fetch("http://127.0.0.1:5000/predict", {
                    method: "POST",
                    headers: {"Content-Type": "application/json"},
                    body: JSON.stringify(payload)
                });

                const data = await response.json();

                if (data.code === 0) {
                    document.getElementById("result").innerHTML = `预测价格:${data.predictions[0].toFixed(2)} 万元`;
                    document.getElementById("error").innerHTML = "";
                } else {
                    document.getElementById("error").innerHTML = `接口错误:${data.error}`;
                }
            } catch (err) {
                document.getElementById("error").innerHTML = `请求失败:${err.message}`;
            }
        });
    </script>
</body>
</html>

访问 http://127.0.0.1:5000/static/index.html ,输入参数,点击“开始预测”,页面会动态显示预测结果。这一步就完成了“机器学习模型 -> Web 接口 -> 前端页面”的全链路打通。

9. 模型评估与优化方向

线性回归虽然是入门算法,但在工程中依然有很多细节需要关注。

9.1 如何评估回归模型

常用指标有三个:

指标 含义 好坏标准
MSE 均方误差,预测值与真实值差的平方均值 越小越好
RMSE 均方根误差,MSE 开平方 越小越好,和原数据同量纲
R2 决定系数,模型解释数据变异的能力 越接近 1 越好,0.9 以上说明拟合良好

train_model.py 中已经输出这三个指标的示例。实际项目里,要重点对比训练集和测试集的指标差异。如果训练集 R2 很高、测试集 R2 很低,说明过拟合。

9.2 数据预处理对结果的影响

现实数据通常不会像示例数据这么干净。需要关注几个问题:

  • 缺失值处理:可以用均值、中位数填充,或者直接删除缺失行。
  • 异常值处理:先用箱线图或 3 倍标准差识别异常值,再决定是删除还是修正。
  • 特征缩放:线性回归对特征量纲不敏感,但如果后续换成带正则化的模型(Ridge、Lasso),就需要做标准化。

9.3 什么时候不能只靠线性回归

如果数据分布表现出明显的非线性,比如价格随面积增长呈曲线上升,线性回归预测偏差会很大。这时可以尝试:

  • 多项式回归:给原始特征增加平方项、交互项。
  • 决策树/随机森林:能处理非线性关系,但可解释性变差。
  • 正则化线性模型:如果特征很多,用 Ridge 或 Lasso 防止过拟合。

9.4 多项式特征扩展示例

from sklearn.preprocessing import PolynomialFeatures

poly = PolynomialFeatures(degree=2, include_bias=False)
X_poly = poly.fit_transform(X)

model_poly = LinearRegression()
model_poly.fit(X_poly, y)

多项式回归的本质仍然是线性回归,只是把原始特征做了组合扩展,让模型能拟合曲线关系。

10. 常见问题与排查方法

问题现象 可能原因 排查方式 解决方案
ModuleNotFoundError: No module named 'sklearn' 未安装 scikit-learn 执行 pip list 查看依赖 执行 pip install scikit-learn
FileNotFoundError: linear_model.pkl 还没有训练模型 检查 models 目录是否存在 先运行 train_model.py
接口返回 400 缺少特征 请求体中字段名不正确 打印请求体,检查 JSON 结构 确保包含 area rooms age 字段
接口返回 500 内部错误 服务端处理异常 查看 Flask 控制台日志 检查传入数据类型是否为数字
Flask 端口被占用 5000 端口已被其它程序使用 执行 `netstat -ano findstr 5000 (Windows) 或 lsof -i :5000` (macOS/Linux)
预测结果偏差很大 训练数据过少或特征不足 查看训练数据分布 增加样本量,补充重要特征
R2 Score 为负数 模型拟合效果极差 检查数据和目标值之间是否存在线性关系 改用非线性模型或增加特征
批量请求耗时较长 单次请求数据量过大 检查预测批次大小 对批量请求做分片,比如每批 1000 条

11. 最佳实践与使用建议

11.1 训练与推理分离

训练脚本和 Web 服务应该分开。训练是离线任务,推理是在线任务。不要每次请求接口时重新训练模型,也不要让训练脚本承担服务职责。

11.2 模型版本管理

训练好的模型文件建议带上版本号:

models/
├── linear_model_v1.0.pkl
├── linear_model_v1.1.pkl
└── linear_model_v2.0.pkl

线上服务加载指定版本,新模型先离线评估,确认指标不下降再切换。

11.3 接口安全与访问控制

Flask 默认监听 127.0.0.1 ,只能本机访问。如果需要在局域网内调用,可以改成 app.run(host="0.0.0.0", port=5000) 。但这意味着同一网络下的所有设备都能访问你的接口,请务必确认使用场景:

  • 不携带敏感数据、不永久保存请求日志。
  • 尽量通过内网部署,不直接暴露公网。
  • 生产环境建议加 API Token 校验。

11.4 建立监控与日志

为接口增加请求日志:

import time

@app.before_request
def before_request():
    request._start_time = time.time()

@app.after_request
def after_request(response):
    duration = time.time() - request._start_time
    print(f"[{time.strftime('%Y-%m-%d %H:%M:%S')}] {request.method} {request.path} - {duration * 1000:.2f}ms")
    return response

这样能直观看到接口响应耗时,后续做性能优化时有数据支撑。

11.5 从小数据开始验证

第一次跑通全流程时,不要用大规模数据。先用几百条样本验证代码正确性,再逐步扩展到完整数据。这能大大减少调试成本。

12. 总结与下一步

本次我们用一套完整示例走通了线性回归的本地训练、模型保存、Flask 接口封装、前端页面调用,是典型的机器学习 Web 化入门路径。 从技术栈来看,你只需要掌握 Python 基础和最简单的 Flask 用法,就能把机器学习模型变成一个可访问的服务 。这也是后续做 AI 应用开发时最常用的基础能力。

最值得先验证的功能 :先跑通 /predict 接口,用一组已知数据人工验证预测结果是否合理。这一步能快速确认模型是否训练正确、特征输入是否匹配。

最容易踩的坑 :请求 JSON 字段名和训练时特征名不一致。Flask 端拿到请求先打印数据格式,能节省大量排查时间。

后续扩展方向

  • 把 sklearn 模型换成 XGBoost 或 LightGBM,对比同一份数据上的精度和性能。
  • 加入特征工程:对真实业务数据做缺失值处理、异常值过滤、特征缩放,观察模型指标变化。
  • 在 Flask 接口中增加数据校验逻辑,保证传入参数非空、类型正确、范围合理。
  • 结合前端做 Web 可视化,把预测结果用图表展示。

线性回归本身并不难,难的是把模型放进真实的 Web 系统里,让它稳定运行、持续提供服务。这篇文章演示的流程可以直接作为 AI 应用开发的学习模板,建议收藏备用。后面再做更复杂的模型时,只要替换训练脚本中的模型部分,接口层和前端层基本不需要大改。

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐