Featured image of post Conditional Neural Processes, CNP

Conditional Neural Processes, CNP

Conditional Neural Processes, CNP

具体流程

image

Step.1 Encoder

$$ \mathbf{r}_i = h_{\theta}(\mathbf{x}_{i}) $$
  • $\mathbf{x}_i \in {X}$
  • $\mathbf{y}_{i} \in Y$
  • $\mathbf{r}_i \in \mathbb{R}^{d}$
  • $h_{\mathbf{\theta}}:X \mapsto \mathbb{R}^{d}$,使用一个神经网络来维护,其参数为$\mathbf{\theta}$

使用编码器意义在于,无论context points输入多少个点,最后都能统一成一个固定长度的向量

但是缺点也很明显,就是丢失原本数据的细节。

Step.2 Aggregate,聚合

$$ \mathbf{r}=\mathbf{r}_1 \oplus \mathbf{r}_2 \oplus \mathbf{r}_3 \oplus \cdots \oplus \mathbf{r}_{N_C} $$
  • $\oplus$:需要满足置换不变性的运算,比如求平均,取最大值
  • $\mathbf{r}_i \in \mathbb{R}^{d}$:经过Encoder的编码

置换不变性的具体说明

$$ Q_{\theta}(f(\mathbf{T})\mid \mathcal{O},\mathbf{T}) = Q_{\theta}(f(\mathbf{T'})\mid \mathcal{O},\mathbf{T'}) = Q_{\theta}(f(\mathbf{T})\mid \mathcal{O'},\mathbf{T}) $$
  • $\mathcal{O}={(\mathbf{x}_i,y_i)}$:已经观察到的数据;
  • $\mathbf{T}=(\mathbf{t}_1,\mathbf{t}_2,\ldots,\mathbf{t}_m)$:希望预测的位置;
  • $f(\mathbf{T})=(f(\mathbf{t}_1),f(\mathbf{t}_2),\ldots,f(\mathbf{t}_m))$;
  • $Q_\theta(f(\mathbf{T})\mid \mathcal{O},\mathbf{T})$:模型根据 $\mathcal{O}$,对目标位置 $\mathbf{T}$ 上函数值给出的概率分布。

例:打乱观测数据,不影响预测

假设我们观察到:

$$ O=\big((1,2),(2,4),(3,6)\big), $$

这些点大致满足 $y=2x$。现在希望预测 $x=4$,所以

$$ T=(4). $$

假设CNP给出的预测分布是

$$ f(4)\mid O,T\sim \mathcal N(8,0.3^2). $$

现在把观测数据的顺序打乱:

$$ O'=\big((3,6),(1,2),(2,4)\big). $$

虽然排列顺序变了,但包含的信息完全相同,因此模型仍然应该给出

$$ f(4)\mid O',T\sim \mathcal N(8,0.3^2). $$

也就是

$$Q_\theta(f(T)\mid O,T)=Q_\theta(f(T)\mid O',T).$$

可以把它类比成计算平均数:

$$ \frac{2+4+6}{3}=\frac{6+2+4}{3}. $$

数字的排列顺序不影响平均数。同样,CNP通常会用求和或求平均的方式汇总观测信息,因此观测点的顺序不会影响最终结果。

$\oplus$的一些方法

  • 取均值:根据加法交换律,这个事情是很显然的
  • 取最值
  • …..

Step.3 Decoder

$$ \phi_j = g_{\theta}(\mathbf{x}^*,\mathbf{r}) $$
  • $g_\theta$ 是解码器:使用神经网络来维护
  • $\mathbf{x}^*$ 是需要预测的目标位置;
  • $\mathbf{r}$ 是全部观测点的总体表示,在Step.2 Aggregate,聚合中得到
  • $\phi_j$ 是目标位置上的概率分布参数,比如说正态分布可以得到$\hat{\mu}{j},\hat{\sigma}{j}$,之后就可以得到$f(\mathbf{x}^\star) \sim \mathcal{N}(\hat{\mu}{j},\hat{\sigma}{j})$

image

训练CNPs

Loss Function

先定义损失函数,CNP原paper中选用了NLL作为损失函数

$$ L(\theta) = -\mathbb{E}_{f \sim D} \left\{\mathbb{E}_N\left[\log Q_{\theta}(\{y_{i}\}_{i=1}^{n}) \mid O_{N} ,\{\mathbf{x}_i\} \right]_{i=1}^{n }\right\} $$
  • $f \sim D$:表示从函数分布(或任务分布)$D$ 中随机抽取一条函数 $f$。

    例如,$D$ 中可能包含很多不同的函数:

    $$ f_1(x)=2x+1,\qquad f_2(x)=-x+3,\qquad f_3(x)=\sin x. $$

    每次训练随机抽取其中一条,用来生成训练数据:

    $$ y_i=f(x_i). $$

    在CNP中,一条函数通常可以看作一个任务。

  • $N \sim U(1,n)_{\text{discrete}}$

  • $Q_{\theta}(y_{i} \mid O, \mathbf{x}_i)$:已知观测集合 $O_N$,模型对位置 $x_i$ 上函数值 $y_i$ 所作的概率预测。

  • $O_N$:表示由 $N$ 个已知观测点组成的观测集合,也叫作上下文集合;或者说是Contents Points 构成的集合

目标是要最小化该函数,即

$$ J(\theta) = \min L(\theta) $$

反向传播

略,pytorch会自动处理的

Gradient Method

image

Gradient Descent

抄袭资料

全文不亚于是抄袭洗稿这两篇文章的,希望支持一下原作者~