基于tensor network的MNIST分类方案咨询与实现疑问
MNIST分类的张量网络策略可行性分析
你的整体思路是可行的,这是典型的基于矩阵乘积态(MPS)的张量网络分类方案,适配MNIST这种高维输入场景,下面逐点拆解合理性:
- 像素映射步骤:
f(p_i)=(p_i, 1-p_i)把每个像素的灰度值编码为二元特征向量,本质是将784维的输入向量转化为784个2维张量的直积空间(对应2^784维),这是MPS处理高维输入的标准预处理方式,能把每个像素的信息转化为张量网络可操作的局部张量。 - 高维数组转MPS:通过逐次SVD分解将高维张量拆分为链式的MPS结构,是MPS的核心构造方法——利用低秩近似压缩高维数据,大幅降低后续计算的复杂度,完全适配784个像素的高维场景。
- 初始化同结构权重W:MPS的分类器本质是用另一个同结构的MPS作为权重,通过两个MPS的收缩计算分类得分,结构一致才能保证对应位置的张量正确匹配收缩,这一步逻辑通顺。
- 收缩得标量做分类:两个MPS的收缩等价于对所有局部张量的对应维度做内积并累积,最终得到的标量可以作为二分类的得分(比如正类得分),是MPS分类的经典思路。
- 损失优化:用交叉熵等损失函数优化W的张量元素,属于标准的监督学习流程,MPS的参数就是各个局部张量的元素,可通过梯度下降等方法更新。
张量收缩(Contraction)原理与手动实现
收缩原理
对于同结构的MPS T 和 W,每个位置的局部张量结构为:T_i[a_{i-1}, s_i, a_i]、W_i[b_{i-1}, s_i, b_i],其中:
a_{i-1}/a_i、b_{i-1}/b_i是MPS的键维度(bond dimension),控制MPS的表达能力和计算复杂度;s_i是像素映射后的2维特征维度(对应p_i和1-p_i)。
收缩的核心逻辑是:
- 对每个位置的
s_i维度做内积(即对应元素相乘后求和),得到局部张量的交互结果; - 将相邻位置的键维度依次连接收缩,最终整个MPS链收缩为一个标量(二分类得分)。
手动实现(无内置收缩函数)
用Python+Numpy手动实现两个MPS的收缩,以下是完整代码:
步骤1:构造示例MPS(模拟MNIST输入的简化版)
import numpy as np # 模拟MNIST的784个像素,这里简化为N=2个像素,键维度D=2 # T的每个张量shape:[左键维度, 2(特征维度), 右键维度] T = [ np.random.rand(1, 2, 2), # 第一个像素:左键=1(链起点),特征=2,右键=2 np.random.rand(2, 2, 1) # 第二个像素:左键=2,特征=2,右键=1(链终点) ] # 同结构的权重MPS W W = [ np.random.rand(1, 2, 2), np.random.rand(2, 2, 1) ]
步骤2:手动实现局部特征维度的收缩
def contract_feature_dim(t, w): """对单个位置的特征维度(2维)做收缩求和""" t_in, _, t_out = t.shape w_in, _, w_out = w.shape # 初始化收缩结果张量 t_w = np.zeros((t_in, t_out, w_in, w_out)) # 遍历特征维度的两个取值 for s in range(2): t_slice = t[:, s, :] # 取t的第s个特征切片,shape=[t_in, t_out] w_slice = w[:, s, :] # 取w的第s个特征切片,shape=[w_in, w_out] # 外积并累加 t_w += t_slice[:, :, np.newaxis, np.newaxis] * w_slice[np.newaxis, np.newaxis, :, :] return t_w
步骤3:手动实现MPS链的逐段收缩
def contract_mps_manual(T, W): """手动实现两个MPS的完整收缩,得到标量得分""" # 初始收缩结果:对应链起点的键维度(1x1) result = np.ones((1, 1)) for t, w in zip(T, W): # 先收缩当前位置的特征维度 t_w = contract_feature_dim(t, w) # 把当前结果和局部收缩结果做键维度的收缩 r_in, r_out = result.shape t_in, t_out, w_in, w_out = t_w.shape new_result = np.zeros((t_out, w_out)) # 手动遍历所有维度做收缩求和 for ri in range(r_in): for ro in range(r_out): new_result += result[ri, ro] * t_w[ro, :, ri, :] result = new_result # 最终得到1x1的标量,取出数值 return result.item()
测试收缩结果
scalar_score = contract_mps_manual(T, W) print("二分类得分:", scalar_score)
扩展到多分类任务的实现
核心概念
给权重MPS W 添加一个标签维度l(l=10对应MNIST的10个数字),让每个局部张量的结构变为 W_i[a_{i-1}, s_i, a_i, l]。此时收缩过程中,标签维度不参与任何收缩操作,最终会保留下来形成一个l维向量,再通过Softmax转换为0-1的概率分布,最大概率对应的索引就是类别标签。
手动实现多分类收缩
def contract_mps_multi_class(T, W): """实现带标签维度的MPS收缩,得到10维概率向量""" l = W[0].shape[-1] # 获取标签数量(MNIST对应10) # 初始结果带标签维度,shape=[1,1,l] result = np.ones((1, 1, l)) for t, w in zip(T, W): t_in, _, t_out = t.shape w_in, _, w_out, _ = w.shape # 初始化局部收缩结果,带标签维度 t_w = np.zeros((t_in, t_out, w_in, w_out, l)) # 遍历特征维度做收缩 for s in range(2): t_slice = t[:, s, :] # [t_in, t_out] w_slice = w[:, s, :, :] # [w_in, w_out, l] t_w += t_slice[:, :, np.newaxis, np.newaxis, np.newaxis] * w_slice[np.newaxis, np.newaxis, :, :, :] # 收缩键维度,保留标签维度 r_in, r_out, _ = result.shape new_result = np.zeros((t_out, w_out, l)) for ri in range(r_in): for ro in range(r_out): for label in range(l): new_result[:, :, label] += result[ri, ro, label] * t_w[ro, :, ri, :, label] result = new_result # 压缩掉多余的键维度,得到l维得分向量 score_vector = result.squeeze() # Softmax转换为0-1的概率分布 prob_vector = np.exp(score_vector) / np.sum(np.exp(score_vector)) return prob_vector
测试多分类收缩
# 构造带标签维度的权重MPS(l=10) W_multi = [ np.random.rand(1, 2, 2, 10), np.random.rand(2, 2, 1, 10) ] prob_vector = contract_mps_multi_class(T, W_multi) print("多分类概率向量:", prob_vector) print("预测类别:", np.argmax(prob_vector))
内容的提问来源于stack exchange,提问作者Mayssa
相关产品推荐
相关产品推荐

