-
Notifications
You must be signed in to change notification settings - Fork 16
Expand file tree
/
Copy pathdemo-scalar.py
More file actions
39 lines (32 loc) · 821 Bytes
/
Copy pathdemo-scalar.py
File metadata and controls
39 lines (32 loc) · 821 Bytes
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
from memo import memo
import jax
import jax.numpy as np
from enum import IntEnum
## Scalar implicature
NN = 10_000
N = np.arange(NN + 1) # number of nice people
class U(IntEnum):
NONE = 0
SOME = 1
ALL = 2
@jax.jit
def meaning(n, u): # (none) (some) (all)
return np.array([ n == 0, n > 0, n == NN ])[u]
@memo
def scalar[n: N, u: U]():
listener: thinks[
speaker: chooses(n in N, wpp=1),
speaker: chooses(u in U, wpp=imagine[
listener: knows(u),
listener: chooses(n in N, wpp=meaning(n, u)),
Pr[listener.n == n]
])
]
listener: observes [speaker.u] is u
listener: chooses(n in N, wpp=E[speaker.n == n])
return Pr[listener.n == n]
scalar() # warm up JIT
import time
t_s = time.time()
scalar()
print(time.time() - t_s)