Skip to main content
QUICK REVIEW

[论文解读] Tensor Networks for Probabilistic Sequence Modeling

Jacob Miller, Guillaume Rabusseau|arXiv (Cornell University)|Mar 2, 2020
Tensor decomposition and applications参考文献 42被引用 5
一句话总结

本文提出了一种用于概率序列建模的统一矩阵乘积态(u-MPS)模型,利用张量网络结构实现序列的 O(log n) 深度并行评估,并提出一种新颖的递归采样算法,可基于任意正则表达式生成条件序列。该方法在数据有限的序列任务中实现了最先进水平的泛化能力,并支持新型结构化生成与正则化形式。

ABSTRACT

Tensor networks are a powerful modeling framework developed for computational many-body physics, which have only recently been applied within machine learning. In this work we utilize a uniform matrix product state (u-MPS) model for probabilistic modeling of sequence data. We first show that u-MPS enable sequence-level parallelism, with length-n sequences able to be evaluated in depth O(log n). We then introduce a novel generative algorithm giving trained u-MPS the ability to efficiently sample from a wide variety of conditional distributions, each one defined by a regular expression. Special cases of this algorithm correspond to autoregressive and fill-in-the-blank sampling, but more complex regular expressions permit the generation of richly structured data in a manner that has no direct analogue in neural generative models. Experiments on sequence modeling with synthetic and real text data show u-MPS outperforming a variety of baselines and effectively generalizing their predictions in the presence of limited data.

研究动机与目标

  • 开发一种基于张量网络的可微分、可微分序列模型,避免使用非线性激活函数。
  • 通过 u-MPS 实现长序列的高度并行评估,达到 O(log n) 的评估深度。
  • 设计一种递归采样算法,可基于任意正则表达式生成序列,扩展了自回归和填空式生成的范畴。
  • 通过正则表达式约束惩罚或激励模式匹配,探索序列建模中的新型正则化技术。
  • 在合成数据集和真实世界文本数据集上,展示模型在泛化能力和结构化生成方面的卓越表现。

提出的方法

  • 该模型使用统一矩阵乘积态(u-MPS)作为无非线性激活函数的可微分序列模型,依赖乘法张量交互作用。
  • u-MPS 通过使用 Adam 优化器在序列数据上进行梯度下降训练,损失函数为负对数似然。
  • 提出一种新颖的递归采样算法 REGSAMP,通过将正则表达式 R 递归分解为子表达式,并从其对应的转移算子中采样,实现从 u-MPS 分布中无偏采样,且条件基于正则表达式 R。
  • 该算法利用 u-MPS 转移算子与正则表达式结构之间的对应关系,即使在处理复杂正则表达式模式(如 Σ*tΣ* 或 R1|R2)时也能实现高效采样。
  • 通过修改训练目标以偏好或惩罚匹配给定正则表达式的序列,该方法支持正则化,适用于偏差缓解和代码生成等应用。
  • 使用 JAX 的 JIT 编译优化递归正则表达式采样实现,尽管其通用性较强,但仍显著降低了计算开销。

实验结果

研究问题

  • RQ1u-MPS 模型能否实现长序列的高效并行评估?此类评估的理论深度复杂度是多少?
  • RQ2u-MPS 模型能否基于任意正则表达式生成序列?其能力如何超越标准的自回归或填空式采样?
  • RQ3当在有限数据上训练时,u-MPS 模型在长序列上的泛化能力如何,特别是在非局部相关性方面?
  • RQ4正则表达式条件采样与正则化能否有效应用于真实世界文本生成任务,如电子邮件地址生成或偏差缓解?
  • RQ5递归采样算法在不同正则表达式结构下的计算与内存效率如何?

主要发现

  • u-MPS 模型以 O(log n) 深度评估长度为 n 的序列,支持高度并行的推理与训练,相比标准 RNN 显著提升。
  • 递归采样算法 REGSAMP 能高效生成任意正则表达式 R 条件下的无偏样本,包括复杂模式如 Σ*tΣ* 或 R1|R2,其时间复杂度为 O(L_R d D^3),内存复杂度为 O(L_R D^2)。
  • 在合成 Tomita 语法数据集上,u-MPS 模型在采样与字符串补全任务中均优于 LSTM 和 Transformer,尤其在训练数据有限(1,000 个 vs. 10,000 个字符串)时表现更优。
  • 该模型成功泛化了非局部相关性(如奇偶性与平衡约束),超越了训练序列长度,展现出强大的归纳偏置。
  • 在真实世界电子邮件地址生成任务中,u-MPS 模型在无条件采样字符串上达到 98.7% 的正确率,判别标准为正则表达式 R_e = [\w-.]+@([\w-]+.)*[\w-][\w-]+。
  • 在 JAX 中使用 JIT 编译显著降低了通用正则表达式采样算法的开销,使其在真实部署中具备可行性。

更好的研究,从现在开始

从阅读论文到最终审阅,大幅缩短您的研究时间。

无需绑定信用卡

本解读由 AI 生成,并经人工编辑审核。