Numpy
numpy的一些基础语法
NumPy 基本语法总结
NumPy 是 Python 中常用的科学计算库,主要用于处理数组、矩阵以及数值计算。在机器学习代码中,经常使用 NumPy 来读取数据、处理数组、计算距离、排序和划分数据集。
一般导入方式为:
1
import numpy as np
这里的 np 是 numpy 的简写,之后就可以通过 np.xxx() 调用 NumPy 中的函数。
创建数组
从列表创建数组
1
a = np.array([1, 2, 3, 4])
结果为:
1
array([1, 2, 3, 4])
二维数组:
1
2
3
4
b = np.array([
[1, 2, 3],
[4, 5, 6]
])
可以理解为矩阵:
1
2
[[1, 2, 3],
[4, 5, 6]]
指定数据类型
1
a = np.array([1.2, 2.5, 3.8], dtype=int)
结果为:
1
array([1, 2, 3])
其中:
1
dtype=int
表示把数组中的元素转换为整数类型。
在 MNIST 代码中:
1
np.array(m_x[0], dtype=int)
表示把第一张图片的像素值转换成整数数组。
查看数组信息
查看数组形状
1
a.shape
例如:
1
2
3
4
5
6
a = np.array([
[1, 2, 3],
[4, 5, 6]
])
print(a.shape)
输出:
1
(2, 3)
表示这个数组有 2 行 3 列。
如果:
1
m_x.shape
输出为:
1
(1000, 784)
表示有 1000 个样本,每个样本有 784 个特征。
对于 MNIST 数据集来说,就是:
1
1000 张图片,每张图片 784 个像素
查看数组长度
1
len(m_x)
表示 m_x 中有多少个样本。
如果:
1
m_x.shape = (1000, 784)
那么:
1
len(m_x)
结果为:
1
1000
查看数据类型
1
a.dtype
例如:
1
2
a = np.array([1, 2, 3])
print(a.dtype)
可能输出:
1
int64
创建特殊数组
创建全 0 数组
1
a = np.zeros(5)
结果为:
1
[0. 0. 0. 0. 0.]
创建长度为 10 的全 0 数组:
1
label_statistic = np.zeros(shape=[10])
结果为:
1
[0. 0. 0. 0. 0. 0. 0. 0. 0. 0.]
在 KNN 代码中,这个数组用于统计 0 到 9 每个类别出现的次数。
创建全 1 数组
1
a = np.ones(5)
结果为:
1
[1. 1. 1. 1. 1.]
创建二维全 1 数组:
1
a = np.ones((2, 3))
结果为:
1
2
[[1. 1. 1.]
[1. 1. 1.]]
创建连续数字数组
1
a = np.arange(10)
结果为:
1
[0 1 2 3 4 5 6 7 8 9]
也可以指定起点、终点和步长:
1
a = np.arange(2, 10, 2)
结果为:
1
[2 4 6 8]
在数据打乱代码中:
1
np.arange(len(m_x))
表示生成所有样本的索引:
1
[0, 1, 2, 3, ..., len(m_x)-1]
改变数组形状
reshape
reshape 用来改变数组的形状。
1
a = np.arange(6)
结果为:
1
[0 1 2 3 4 5]
将其改成 2 行 3 列:
1
b = np.reshape(a, [2, 3])
结果为:
1
2
[[0 1 2]
[3 4 5]]
也可以写成:
1
b = a.reshape(2, 3)
在 MNIST 代码中:
1
data = np.reshape(np.array(m_x[0], dtype=int), [28, 28])
表示把一张图片的 784 个像素,重新变成 28 × 28 的图像矩阵。
因为:
1
28 × 28 = 784
所以可以把一维数组:
1
(784,)
变成二维矩阵:
1
(28, 28)
数组索引和切片
一维数组索引
1
a = np.array([10, 20, 30, 40])
取第一个元素:
1
a[0]
结果为:
1
10
取第二个元素:
1
a[1]
结果为:
1
20
Python 的索引从 0 开始。
二维数组索引
1
2
3
4
a = np.array([
[1, 2, 3],
[4, 5, 6]
])
取第 1 行:
1
a[0]
结果为:
1
[1, 2, 3]
取第 2 行第 3 列:
1
a[1, 2]
结果为:
1
6
也可以写成:
1
a[1][2]
切片
1
a = np.array([0, 1, 2, 3, 4, 5])
取前 3 个元素:
1
a[:3]
结果为:
1
[0 1 2]
从下标 3 开始取到最后:
1
a[3:]
结果为:
1
[3 4 5]
取中间一段:
1
a[1:4]
结果为:
1
[1 2 3]
注意:
1
a[1:4]
包含下标 1,不包含下标 4。
在训练集和测试集划分代码中:
1
x_train, x_test = m_x[:split], m_x[split:]
含义是:
1
2
m_x[:split] → 前 80% 作为训练集
m_x[split:] → 后 20% 作为测试集
用索引数组重新排列数据
在数据打乱代码中:
1
2
3
idx = np.random.permutation(np.arange(len(m_x)))
m_x = m_x[idx]
m_y = m_y[idx]
假设:
1
idx = [2, 0, 3, 1]
那么:
1
m_x[idx]
表示按照下面的顺序重新取数据:
1
2
3
4
先取 m_x[2]
再取 m_x[0]
再取 m_x[3]
再取 m_x[1]
这样就实现了随机打乱数据。
关键是:
1
2
m_x = m_x[idx]
m_y = m_y[idx]
图片数据和标签数据必须使用同一个 idx 打乱,这样才能保证图片和标签仍然一一对应。
不能分别对 m_x 和 m_y 单独随机打乱,否则图片和标签会错位。
随机相关语法
设置随机种子
1
np.random.seed(0)
作用是固定随机结果。
如果没有这句,每次运行:
1
np.random.permutation(...)
得到的随机顺序可能不同。
有了:
1
np.random.seed(0)
每次运行得到的随机顺序就一样,方便复现实验结果。
随机打乱数组
1
idx = np.random.permutation(np.arange(10))
可能得到:
1
[2 8 4 9 1 6 7 3 0 5]
它会返回一个被随机打乱后的新数组。
在 KNN 代码中,它用于生成随机索引,从而打乱数据集顺序。
数学运算
NumPy 可以直接对数组进行数学运算。
1
2
a = np.array([1, 2, 3])
b = np.array([4, 5, 6])
加法:
1
a + b
结果为:
1
[5 7 9]
减法:
1
a - b
结果为:
1
[-3 -3 -3]
乘法:
1
a * b
结果为:
1
[ 4 10 18]
平方:
1
np.square(a)
结果为:
1
[1 4 9]
开方:
1
np.sqrt(a)
求和:
1
np.sum(a)
结果为:
1
6
欧氏距离计算
在 KNN 代码中,距离函数是:
1
2
def distance(a, b):
return np.sqrt(np.sum(np.square(a - b)))
它的计算过程是:
1
两个数组相减 → 每一项平方 → 所有平方值求和 → 开方
例如两个样本:
1
2
a = np.array([1, 2])
b = np.array([4, 6])
计算过程为:
1
2
a - b
# [-3, -4]
1
2
np.square(a - b)
# [9, 16]
1
2
np.sum(np.square(a - b))
# 25
1
2
np.sqrt(np.sum(np.square(a - b)))
# 5
所以两个点之间的距离为 5。
排序相关语法
np.sort
1
2
a = np.array([30, 10, 20])
np.sort(a)
结果为:
1
[10 20 30]
np.sort() 返回排序后的值。
np.argsort
1
2
a = np.array([30, 10, 20])
np.argsort(a)
结果为:
1
[1 2 0]
np.argsort() 返回的是排序后的下标,而不是排序后的值。
解释如下:
1
2
3
最小值 10 原来的下标是 1
第二小 20 原来的下标是 2
最大值 30 原来的下标是 0
在 KNN 代码中:
1
knn_indices = np.argsort(dis)
表示按照距离从小到大排序,并返回对应训练样本的下标。
找最大值位置
np.argmax
1
2
a = np.array([1, 5, 3])
np.argmax(a)
结果为:
1
1
因为最大值是 5,它的下标是 1。
在 KNN 代码中:
1
return np.argmax(label_statistic)
表示返回出现次数最多的类别。
例如:
1
label_statistic = [0, 1, 0, 0, 0, 0, 0, 3, 0, 0]
最大值是 3,对应下标是 7,所以预测类别为:
1
7
读取文本数据
np.loadtxt
1
2
m_x = np.loadtxt('mnist_x', delimiter=' ')
m_y = np.loadtxt('mnist_y')
np.loadtxt() 用于从文本文件中读取数据。
其中:
1
delimiter=' '
表示同一行中的数据用空格分隔。
对于 mnist_x:
1
2
3
4
一行是一张图片
一行中有 784 个像素值
像素值之间用空格隔开
不同图片之间用换行分隔
对于 mnist_y:
1
2
一行是一个标签
每个标签对应 mnist_x 中的一张图片
NumPy 常用函数总结
| 语法 | 作用 |
|---|---|
np.array() | 创建 NumPy 数组 |
np.zeros() | 创建全 0 数组 |
np.ones() | 创建全 1 数组 |
np.arange() | 创建连续数字数组 |
np.reshape() | 改变数组形状 |
np.loadtxt() | 从文本文件读取数据 |
np.random.seed() | 固定随机结果 |
np.random.permutation() | 随机打乱数组 |
np.sum() | 求和 |
np.square() | 平方 |
np.sqrt() | 开方 |
np.sort() | 返回排序后的值 |
np.argsort() | 返回排序后的下标 |
np.argmax() | 返回最大值所在下标 |
结合 KNN 代码理解
读取图片和标签
1
2
m_x = np.loadtxt('mnist_x', delimiter=' ')
m_y = np.loadtxt('mnist_y')
作用:
1
2
m_x:读取图片像素数据
m_y:读取图片对应的标签
显示第一张图片
1
data = np.reshape(np.array(m_x[0], dtype=int), [28, 28])
作用:
1
2
3
取出第一张图片
将像素值转成整数
把 784 个像素还原成 28×28 的二维图像矩阵
打乱数据集
1
2
3
4
np.random.seed(0)
idx = np.random.permutation(np.arange(len(m_x)))
m_x = m_x[idx]
m_y = m_y[idx]
作用:
1
2
3
生成随机索引
用同一个索引打乱图片和标签
保证图片和标签仍然一一对应
划分训练集和测试集
1
2
3
4
5
ratio = 0.8
split = int(len(m_x) * ratio)
x_train, x_test = m_x[:split], m_x[split:]
y_train, y_test = m_y[:split], m_y[split:]
作用:
1
2
前 80% 数据作为训练集
后 20% 数据作为测试集
计算两个样本之间的距离
1
2
def distance(a, b):
return np.sqrt(np.sum(np.square(a - b)))
作用:
1
2
计算两个样本之间的欧氏距离
距离越小,说明两个样本越相似
找最近的 K 个样本
1
2
knn_indices = np.argsort(dis)
knn_indices = knn_indices[:self.k]
作用:
1
2
按照距离从小到大排序
取出距离最近的 K 个训练样本的下标
统计类别并返回预测结果
1
2
3
4
5
6
7
label_statistic = np.zeros(shape=[self.label_num])
for index in knn_indices:
label = int(self.y_train[index])
label_statistic[label] += 1
return np.argmax(label_statistic)
作用:
1
2
统计 K 个近邻中每个类别出现的次数
返回出现次数最多的类别
总结
NumPy 的核心作用是高效处理数组。
在这个 KNN 项目中,NumPy 主要完成了以下任务:
1
2
3
4
5
6
7
8
9
1. 读取数据
2. 表示图片像素矩阵
3. 改变数组形状
4. 打乱数据集
5. 划分训练集和测试集
6. 计算样本之间的距离
7. 对距离进行排序
8. 统计类别票数
9. 得到最终预测结果
对于机器学习入门来说,最常用的 NumPy 语法包括:
1
2
3
4
5
6
7
8
9
10
11
12
np.array()
np.reshape()
np.zeros()
np.arange()
np.random.seed()
np.random.permutation()
np.sum()
np.square()
np.sqrt()
np.argsort()
np.argmax()
np.loadtxt()