aboutsummaryrefslogtreecommitdiff
path: root/test.py
blob: d1db16dd5bc15ce645d0ecb4ec009c1f91842688 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
import numpy as np
import matplotlib.pyplot as plt

A = np.random.uniform(0.1, 0.9)
B = np.random.uniform(0.1, 0.9)

def q_xt_xt_1_simple(x, t):
    mean = A * x
    std = B
    return np.random.normal(mean, std)

def q_xt_x0_simple_damien(x, t):
    mean = A ** t * x
    std = np.sqrt(sum([A ** (2*i) for i in range(t)])) * B
    return np.random.normal(mean, std)


def q_xt_xt_1(x, t):
    alpha = ALPHA[t]
    mean = np.sqrt(alpha) * x
    std = 1 - alpha
    return np.random.normal(mean, std)

def q_xt_x0_paper(x, t):
    alpha_bar = ALPHA_BAR[t]
    mean = np.sqrt(alpha_bar) * x
    std = 1 - alpha_bar
    return np.random.normal(mean, std)

def q_xt_x0_damien(x, t):
    alpha_bar = ALPHA_BAR[t]
    cum_sq_sum = sum([np.prod(ALPHA[s+2:t+1]) * (1 - ALPHA[s+1])**2 for s in range(t)])
    mean = np.sqrt(alpha_bar) * x
    std = np.sqrt(cum_sq_sum)
    return np.random.normal(mean, std)

T = 100
BETA = np.concatenate(([0], np.linspace(1e-4, 2e-2, T)))
ALPHA = 1 - BETA
ALPHA_BAR = np.cumprod(ALPHA)

N = int(1e6)
x0 = 1

xs_implicit = np.array([x0] * N)
for t in range(1, T+1):
    xs_implicit = q_xt_xt_1(xs_implicit, t)

xs_explicit_paper = q_xt_x0_paper(np.array([x0] * N), T)
xs_explicit_damien = q_xt_x0_damien(np.array([x0] * N), T)

plt.figure()
bins = np.linspace(min(
                xs_implicit.min(),
                xs_explicit_paper.min(),
                xs_explicit_damien.min(),
            ), max(
                xs_implicit.max(),
                xs_explicit_paper.max(),
                xs_explicit_damien.max(),
            ), 100)
plt.hist(xs_implicit, bins=bins, alpha=0.5, label="q_xt_xt_1")
plt.hist(xs_explicit_paper, bins=bins, alpha=0.5, label="q_xt_x0_paper")
plt.hist(xs_explicit_damien, bins=bins, alpha=0.5, label="q_xt_x0_damien")
plt.legend()
plt.savefig("q_xt_xt_1_vs_q_xt_x0.tmp.png")
plt.show()

xs_implicit = np.array([x0] * N)
for t in range(1, T+1):
    xs_implicit = q_xt_xt_1_simple(xs_implicit, t)

xs_explicit_damien = q_xt_x0_simple_damien(np.array([x0] * N), T)

plt.figure()
bins = np.linspace(min(xs_implicit.min(), xs_explicit_damien.min()), max(xs_implicit.max(), xs_explicit_damien.max()), 100)
plt.hist(xs_implicit, bins=bins, alpha=0.5, label="q_xt_xt_1_simple")
plt.hist(xs_explicit_damien, bins=bins, alpha=0.5, label="q_xt_x0_simple_damien")
plt.legend()
plt.savefig("q_xt_xt_1_simple_vs_q_xt_x0_simple.tmp.png")
plt.show()

A1, B1, A2, B2, A3, B3 = np.random.uniform(0, 1, 6)
x1 = np.random.normal(A1, B1, N)
x2 = np.random.normal(A2 * x1, B2, N)
x3 = np.random.normal(A3 * x2, B3, N)
x3_ = np.random.normal(A1 * A2 * A3, np.sqrt(A3**2 * A2**2 * B1**2 + A3**2 * B2**2 + B3**2), N)

plt.figure()
bins = np.linspace(min(x3.min(), x3_.min()), max(x3.max(), x3_.max()), 100)
plt.hist(x3, bins=bins, alpha=0.5, label="normal")
plt.hist(x3_, bins=bins, alpha=0.5, label="product")
plt.legend()
plt.savefig("product_normal.tmp.png")
plt.show()