机器学习中的表格数据分类与回归实战

1. 表格数据分类

在机器学习中,表格数据分类是一个重要的任务。分类报告包含了每个类别的详细指标,如精度、召回率、F1 分数,整体准确率以及混淆矩阵。这些信息可用于评估分类器的性能,并比较不同分类器的结果。

为了使特征工程和分类过程更易于管理,并能将工作流程导出到 C++,可以创建一个分类链。分类链本质上是一系列按顺序对数据进行操作的列表,每个操作的输出作为下一个操作的输入,直到生成最终结果。以下是创建水果数据集完整分类链的示例代码:

from tinyml4all.tabular.classification import Table, Chain
from tinyml4all.tabular.classification.models import RandomForest
from tinyml4all.tabular.features import Scale

table = Table.read_csv("fruits.csv")
chain = Chain(
    Scale("minmax"),
    RandomForest()
)
classified = chain(table)

第一次在表格上调用链时,它会遍历每个步骤并进行拟合,以学习其内部参数。初始训练阶段后,可将其应用于其他表格,此时链会记住参数,仅对新表格应用相应的转换。示例代码如下:

table1 = Table.read_csv("dataset_1.csv")
table2 = Table.read_csv("dataset_2.csv")
chain = Chain(
    Scale("minmax"),
    RandomForest()
)
# train chain on table1
chain(table1)
# apply chain to table2
# this will use the parameters learned on table1
# e.g., use min/max from table1
table2_classified = chain(table2)

完成机器学习工作流程后,可将管道转换为 C++ 代码,以便导入到嵌入式项目中。示例代码如下:

chain.convert_to("c++", class_name="FruitChain", save_to="FruitChain.h")

生成的代码是一个独立的类,可导入到任何 C++ 项目中。以下是在 Arduino 草图中使用该代码进行水果分类的示例:

/**
 * Listing 2-25
 * Predict fruit from RGB color components.
 *
 * Required hardware: Arduino Nano BLE Sense.
 */
#include <Arduino_APDS9960.h>
#include <tinyml4all.h>
// put the generated file from Python inside the Arduino
// sketch folder!
#include "./FruitChain.h"

tinyml4all::APDS9960 sensor;
tinyml4all::FruitChain chain;

void setup() {
    Serial.begin(115200);
    while (!Serial);
    Serial.println("Fruits classification example");
    sensor.begin();
}

void loop() {
    sensor.readColor();
    // chain(input) will return true on success
    // false on error
    if (!chain(sensor.r, sensor.g, sensor.b))
        return;
    // the predicted human-readable label is in
    // chain.output.classification.label
    // the numeric output (0, 1, 2, ...) is in
    // chain.output.classification.idx
    Serial.print("I think this is ");
    Serial.println(chain.output.classification.label);
    delay(1000);
}

需要注意的是,分类器只能识别训练期间见过的案例。如果有“无感兴趣对象”的情况,需要收集相关数据。

2. 表格数据回归

表格数据回归与分类有许多相同的步骤,如数据捕获和特征工程,但在数据绘图和算法选择上有所不同。以通过 RGB 颜色分量推断距离的接近度计项目为例,其步骤如下:
1. 收集数据(RGB + 距离)
2. 加载和检查数据,发现异常和不良数据,理解输入和输出之间的关系
3. 进行特征工程,提高建模效果
4. 训练机器学习模型进行回归任务并评估其性能
5. 将整个回归链转换为 C++ 代码并部署到开发板

2.1 所需硬件

除了颜色传感器,还需要一个距离传感器,如超声波、飞行时间、红外等。Arduino Nano BLE Sense 板内置了接近传感器,但在近距离(<5 cm)效果不佳,因此可使用外部超声波传感器。

2.2 捕获数据

使用 Python 脚本和 Arduino 草图收集 RGB + 距离数据。以下是 Arduino 草图示例:

/**
 * Listing 3-1
 * Collect RGB + distance data in CSV format (unattended).
 *
 * Required hardware: Arduino Nano BLE Sense.
 * Required hardware: Ultrasonic distance sensor (HC-SR04)
 */
#include <Arduino_APDS9960.h>
#include <tinyml4all.h>

#define ECHO 2
#define TRIG 3

using tinyml4all::printCSV;

tinyml4all::APDS9960 sensor;
tinyml4all::Ultrasonic ultrasonic(ECHO, TRIG);

void setup() {
  Serial.begin(115200);
  while (!Serial);
  Serial.println("Collect RGB + distance as CSV");
  // init color and distance sensors
  sensor.begin();
  ultrasonic.begin();
}

void loop() {
  // read R, G, B
  sensor.readColor();
  // read distance in millimeters
  uint16_t distance = ultrasonic.millimiters();
  // print data as CSV
  printCSV(sensor.r, sensor.g, sensor.b, distance);
  delay(1000);
}

Python 脚本示例如下:

from tinyml4all.tabular import capture_serial

capture_serial(
    # * is a wildcard character that matches anything
    # on Windows, this will look like COM1 or similar
    port="/dev/cu.usb*",
    baudrate=115200,
    # file where data will be saved
    save_to="rgb-distance.csv",
    # name of columns
    headings="r,g,b,distance",
    num_samples=100
)

运行代码时,从距离明亮平坦表面 2 - 3 cm 处开始,缓慢移动开发板/传感器。完成后,确保有一个包含 100 行数据的 rgb-distance.csv 文件。

2.3 加载和检查数据

使用以下代码加载单个 CSV 文件:

# note that this time the module is called regression
# instead of classification!
from tinyml4all.tabular.regression import Table

table = Table.read_csv("rgb-distance.csv")
table.set_targets(column="distance")
print(table.describe())

加载表格的代码与分类类似,但导入模块从 tinyml4all.tabular.classification 变为 tinyml4all.tabular.regression set_targets 函数形式未变,但底层逻辑不同,回归任务的目标现在表示连续值。

2.4 绘制回归数据

回归中的彩色散点图和配对图没有意义,因为回归没有类别,而是连续输出。
- 一个输入变量 :可使用以下代码绘制输入变量与输出变量的散点图:

# only plot a single column
table.scatter(column="r")
  • 多个输入变量 :可绘制目标变量与所有输入变量的散点图:
# you can stack the plots horizontally or vertically
# default is vertically
table.scatter(orientation="vertical")

通过观察散点图,如果点与红线距离不远且呈单调增长分布,说明颜色分量与距离之间有很强的关系。

2.5 特征工程

许多回归模型在输入和输出之间存在线性关系时效果较好。但并非所有数据都具有这种关系,可能需要进行单调函数映射。常见的映射包括:
- 2 次方、3 次方、平方根、倒数
- 指数和对数
- Box - Cox 和 Yeo - Johnson

以下是应用单调函数变换的示例代码:

from tinyml4all.tabular.features import Monotonic

# apply the square and cube mapping to all columns
square_and_cube = Monotonic(functions="square, cube")
# apply all the mappings only to the "r" column
only_r = Monotonic(columns="r")
# apply all mappings to all columns
monotonic = Monotonic()
# run the transform on the table
table2 = monotonic(table)
print(table2.describe())

应根据数据选择合适的映射,可通过观察回归图将点的分布形状与可用函数关联起来。

机器学习中的表格数据分类与回归实战(下半部分)

3. 单调函数映射的选择与应用

在回归任务的特征工程中,单调函数映射起着关键作用。不同的数据分布需要不同的映射来建立输入与输出之间的线性关系。下面我们详细探讨如何选择合适的单调函数映射。

3.1 常见单调函数映射的特点
映射类型 特点 适用场景
2 次方、3 次方 增强数据的非线性特征,使数据的变化幅度增大 当数据的变化趋势呈现出二次或三次曲线关系时适用
平方根 压缩数据的变化幅度,使较大值的影响相对减小 数据存在较大的波动,且希望减小极端值的影响时使用
倒数 反转数据的大小顺序,将大值变为小值,小值变为大值 当数据的关系呈现反比例关系时适用
指数 快速增长或衰减数据,强调数据的变化速度 数据的增长或衰减速度随输入变化而变化时使用
对数 压缩数据的动态范围,使数据更易于处理 数据的变化范围非常大,需要将其压缩到一个较小的区间时适用
Box - Cox 和 Yeo - Johnson 更复杂的变换,可根据数据自动调整变换参数 当难以确定具体的映射类型时,可尝试这两种方法
3.2 选择合适映射的方法

通过观察回归图中数据点的分布形状,可以初步判断适合的映射类型。例如,如果数据点呈现出类似二次曲线的分布,那么 2 次方映射可能是一个不错的选择;如果数据点的分布类似于指数增长,那么指数映射可能更合适。

以下是一个简单的决策流程图:

graph TD;
    A[观察回归图] --> B{数据点分布形状};
    B --> C[线性分布] --> D[无需映射];
    B --> E[二次曲线分布] --> F[使用 2 次方映射];
    B --> G[三次曲线分布] --> H[使用 3 次方映射];
    B --> I[指数增长分布] --> J[使用指数映射];
    B --> K[对数分布] --> L[使用对数映射];
    B --> M[反比例分布] --> N[使用倒数映射];
    B --> O[其他复杂分布] --> P[尝试 Box - Cox 或 Yeo - Johnson 映射];
4. 回归模型的训练与评估

在完成数据的特征工程后,就可以进行回归模型的训练和评估了。

4.1 选择合适的回归模型

常见的回归模型有线性回归、决策树回归、随机森林回归等。不同的模型适用于不同的数据特点和问题场景。
| 模型名称 | 特点 | 适用场景 |
| ---- | ---- | ---- |
| 线性回归 | 简单易懂,计算速度快,但只能处理线性关系 | 数据呈现明显的线性关系时使用 |
| 决策树回归 | 可以处理非线性关系,能够自动选择重要特征 | 数据关系复杂,需要进行特征选择时适用 |
| 随机森林回归 | 由多个决策树组成,具有较好的泛化能力 | 数据量较大,且希望模型具有较高的准确性和稳定性时使用 |

4.2 训练回归模型

以随机森林回归为例,以下是训练模型的示例代码:

from tinyml4all.tabular.regression import RandomForestRegressor
from tinyml4all.tabular.regression import Table

table = Table.read_csv("rgb-distance.csv")
table.set_targets(column="distance")

# 划分训练集和测试集
train_table, test_table = table.split(0.8)

model = RandomForestRegressor()
model.fit(train_table)

# 在测试集上进行预测
predictions = model.predict(test_table)
4.3 评估回归模型

评估回归模型的常用指标有均方误差(MSE)、均方根误差(RMSE)和决定系数(R²)。以下是计算这些指标的示例代码:

from sklearn.metrics import mean_squared_error, r2_score
import numpy as np

# 计算均方误差
mse = mean_squared_error(test_table.get_targets(), predictions)
# 计算均方根误差
rmse = np.sqrt(mse)
# 计算决定系数
r2 = r2_score(test_table.get_targets(), predictions)

print(f"均方误差: {mse}")
print(f"均方根误差: {rmse}")
print(f"决定系数: {r2}")
5. 将回归链转换为 C++ 并部署到开发板

完成回归模型的训练和评估后,最后一步是将整个回归链转换为 C++ 代码,并部署到开发板上。

5.1 转换为 C++ 代码

使用以下代码将回归链转换为 C++ 代码:

from tinyml4all.tabular.regression import Chain
from tinyml4all.tabular.features import Monotonic
from tinyml4all.tabular.regression.models import RandomForestRegressor

table = Table.read_csv("rgb-distance.csv")
table.set_targets(column="distance")

chain = Chain(
    Monotonic(),
    RandomForestRegressor()
)

chain.fit(table)
chain.convert_to("c++", class_name="DistanceChain", save_to="DistanceChain.h")
5.2 部署到开发板

以下是在 Arduino 草图中使用转换后的 C++ 代码进行距离预测的示例:

/**
 * Predict distance from RGB color components.
 *
 * Required hardware: Arduino Nano BLE Sense.
 */
#include <Arduino_APDS9960.h>
#include <tinyml4all.h>
// put the generated file from Python inside the Arduino
// sketch folder!
#include "./DistanceChain.h"

tinyml4all::APDS9960 sensor;
tinyml4all::DistanceChain chain;

void setup() {
    Serial.begin(115200);
    while (!Serial);
    Serial.println("Distance prediction example");
    sensor.begin();
}

void loop() {
    sensor.readColor();
    // chain(input) will return true on success
    // false on error
    if (!chain(sensor.r, sensor.g, sensor.b))
        return;
    // the predicted distance is in
    // chain.output.regression.value
    Serial.print("The predicted distance is ");
    Serial.println(chain.output.regression.value);
    delay(1000);
}

通过以上步骤,我们完成了从数据收集、特征工程、模型训练到代码部署的整个机器学习流程,实现了通过 RGB 颜色分量推断距离的功能。在实际应用中,可根据具体需求调整模型和参数,以获得更好的性能。

Logo

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

更多推荐