DiScoFormer:用一个 Transformer 同时学习密度与分数函数

2026-06-30 37 预计阅读时间: 1 分钟
来源: huggingface.co AI 摘要 Original link

Disclaimer: This article is an AI-assisted summary. Read it together with the original source when precision matters. The summary may omit context, version differences, or edge cases and is not official documentation.

预计阅读时间:10 分钟

DiScoFormer 这个标题里有两个很重的词:density(概率密度)和 score(分数函数,通常指对数密度的梯度)。如果一个 Transformer 能在不同分布之间同时处理这两类对象,它瞄准的就不是“又一个序列模型”,而是更底层的概率建模接口:既能估计一个点有多可能出现,也能告诉你应该往哪个方向移动才能更接近高概率区域。

由于这里没有更多论文细节,下面会把“DiScoFormer”当作一个研究方向来解读:一个统一模型如何面向多种分布,联合学习密度与 score。代码部分是可实践的最小化示例,用来帮助你理解这类模型在工程上可能如何组织,而不是声称复现了原论文。

为什么 density 和 score 值得放在一起

在概率建模里,密度函数回答的是:

给定样本 x,它在当前分布下有多大概率?

score 函数回答的是:

如果想让 x 更像来自这个分布,应该往哪个方向调整?

形式上,如果概率密度是 p(x),score 通常写作:

score(x) = ∇x log p(x)

这两个量天然相关。密度给出“地形高度”,score 给出“上坡方向”。在扩散模型、能量模型、分布迁移、采样算法里,score 往往比密度本身更直接可用;而密度又能提供归一化、似然评估和异常检测等能力。

把二者放进同一个模型,有几个潜在好处:

  • 共享表示:模型不用分别为密度估计和 score 估计学习两套特征。
  • 互相约束:score 应该与密度的梯度一致,这给训练带来额外结构。
  • 跨分布泛化:如果输入里显式编码“这是哪一个分布”,模型可能学到分布族之间的共性。
  • 统一接口:下游系统可以用同一个模型做 likelihood、denoising、sampling 或异常分数。

“Across distributions” 难在哪里

单一分布的密度估计已经不简单。跨分布则更像是在问:模型能不能看懂“分布本身”这个对象?

举个例子,二维高斯、混合高斯、环形分布、多峰分布,它们的样本都可以是二维点,但背后的概率地形完全不同。如果只把 x 喂给模型,模型不知道当前要解释的是哪张地形。工程上通常需要额外上下文,例如:

  • 分布参数:均值、协方差、温度、噪声强度等。
  • 样本集合:给模型一批来自目标分布的上下文样本。
  • 任务标签:告诉模型当前是哪个数据域或哪个分布族。
  • 时间步/噪声级别:扩散模型中常见的条件变量。

Transformer 适合做这件事的原因在于:它可以把“查询点 x”和“分布上下文”都看作 token,然后通过注意力机制建模它们之间的关系。这比手写一堆针对特定分布的特征更灵活。

可以这样实践:一个最小 PyTorch 原型

下面的示例不是 DiScoFormer 论文实现,而是一个可运行的玩具项目:用一个小 Transformer 同时预测二维高斯混合分布的 log density 和 score。它展示了这类模型最核心的工程形状:输入点、分布上下文、两个输出头。

运行前需要安装 PyTorch:

pip install torch

保存为 toy_discoformer.py 后运行:

import math
import torch
import torch.nn as nn
import torch.nn.functional as F


torch.manual_seed(7)


def gaussian_mixture_logp_and_score(x, means, sigma=0.45):
    # x: [batch, 2], means: [components, 2]
    diff = x[:, None, :] - means[None, :, :]
    log_each = -0.5 * (diff.square().sum(-1) / sigma**2) - math.log(2 * math.pi * sigma**2)
    logp = torch.logsumexp(log_each - math.log(means.shape[0]), dim=1)

    weights = torch.softmax(log_each, dim=1)
    component_scores = -diff / sigma**2
    score = (weights[:, :, None] * component_scores).sum(dim=1)
    return logp[:, None], score


class TinyDiscoFormer(nn.Module):
    def __init__(self, d_model=64):
        super().__init__()
        self.point_proj = nn.Linear(2, d_model)
        self.context_proj = nn.Linear(2, d_model)
        layer = nn.TransformerEncoderLayer(
            d_model=d_model,
            nhead=4,
            dim_feedforward=128,
            batch_first=True,
        )
        self.encoder = nn.TransformerEncoder(layer, num_layers=2)
        self.logp_head = nn.Linear(d_model, 1)
        self.score_head = nn.Linear(d_model, 2)

    def forward(self, x, context_means):
        # token 0 is the query point; following tokens describe the distribution.
        point_token = self.point_proj(x).unsqueeze(1)
        context_tokens = self.context_proj(context_means).unsqueeze(0).expand(x.shape[0], -1, -1)
        tokens = torch.cat([point_token, context_tokens], dim=1)
        encoded = self.encoder(tokens)[:, 0]
        return self.logp_head(encoded), self.score_head(encoded)


def main():
    means = torch.tensor([[-1.2, -0.8], [1.0, 0.7], [-0.2, 1.4]])
    model = TinyDiscoFormer()
    optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)

    for step in range(800):
        x = torch.empty(256, 2).uniform_(-3, 3)
        target_logp, target_score = gaussian_mixture_logp_and_score(x, means)
        pred_logp, pred_score = model(x, means)
        loss = F.mse_loss(pred_logp, target_logp) + F.mse_loss(pred_score, target_score)

        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        if step % 200 == 0:
            print(f"step={step:04d} loss={loss.item():.4f}")

    test_x = torch.tensor([[0.0, 0.0], [1.2, 0.8], [-2.0, -2.0]])
    with torch.no_grad():
        logp, score = model(test_x, means)
    print("\nquery points:", test_x)
    print("predicted log density:", logp.squeeze(-1))
    print("predicted score:", score)


if __name__ == "__main__":
    main()

这个例子里,context_means 就是“分布上下文”。真实研究系统会复杂得多:上下文可能是一组样本、条件文本、噪声日程,或者可学习的分布 token;训练目标也可能包括去噪 score matching、最大似然、flow matching 或多任务损失。

工程上要盯住的三个边界

联合学习 density 和 score 听起来优雅,但落地时有几个坑要提前看清。

第一,梯度一致性不是自动成立的。 如果模型有一个 log density head 和一个 score head,它们可能学出彼此矛盾的结果。更严格的做法是让 score 直接来自 ∇x logp(x),或者在损失里加入一致性约束。但这样会增加二阶梯度、显存和训练成本。

第二,跨分布泛化取决于上下文设计。 Transformer 并不会凭空理解“分布”。如果上下文 token 没有携带足够信息,模型只是在平均多个任务,表现会变钝。样本数量、上下文编码方式、分布族覆盖范围都会影响泛化。

第三,密度估计很容易被维度惩罚。 在高维数据上,likelihood 与人类感知质量、生成质量、异常程度之间未必一致。工程系统不要只看一个 log likelihood 指标,最好同时观察采样质量、校准误差、下游任务收益和鲁棒性。

什么时候值得关注这类统一模型

如果你的系统只需要分类或普通回归,DiScoFormer 这类方向可能显得过重。但在下面这些场景里,统一 density 与 score 的模型很有吸引力:

  • 需要在多个数据分布之间快速适配,例如不同传感器、不同用户群、不同实验条件。
  • 既要评估样本概率,又要生成、修复或去噪样本。
  • 希望把概率建模能力封装成统一服务,而不是维护多个专用模型。
  • 研究问题本身涉及分布族、采样动力学或 score-based 方法。

采用时可以从小规模验证开始:先选低维或可解析分布,确认 density 与 score 都能学对;再换成真实数据上下文;最后才考虑大模型和复杂训练目标。把它当成概率建模基础设施,而不是单纯的 Transformer 变体,会更容易判断它是否真的适合你的系统。


相关推荐