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

基于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)。

收缩的核心逻辑是:

  1. 对每个位置的s_i维度做内积(即对应元素相乘后求和),得到局部张量的交互结果;
  2. 将相邻位置的键维度依次连接收缩,最终整个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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 20:50:54