-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathlogistic_regression_model.py
More file actions
116 lines (87 loc) · 3.08 KB
/
Copy pathlogistic_regression_model.py
File metadata and controls
116 lines (87 loc) · 3.08 KB
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
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
import numpy as np
import jax.numpy as jnp
from jax import jit, grad, value_and_grad, vmap, random
from jax.scipy.special import logsumexp
import warnings
#ignore by GPU/TPU message (generated by jax module)
warnings.filterwarnings("ignore", message='No GPU/TPU found, falling back to CPU.')
def genCovMat(key, d):
return jnp.eye(d)
def logistic(theta, x):
return 1/(1+jnp.exp(-jnp.dot(theta, x)))
batch_logistic = jit(vmap(logistic, in_axes=(None, 0)))
batch_benoulli = vmap(random.bernoulli, in_axes=(0, 0))
def gen_data(key, dim, N):
"""
Generate data with dimension `dim` and `N` data points
Parameters
----------
key: uint32
random key
dim: int
dimension of data
N: int
Size of dataset
Returns
-------
theta_true: ndarray
Theta array used to generate data
X: ndarray
Input data, shape=(N,dim)
y_data: ndarray
Output data: 0 or 1s. shape=(N,)
"""
key, subkey1, subkey2, subkey3 = random.split(key, 4)
print(f"generating data, with N={N} and dim={dim}")
theta_true = random.normal(subkey1, shape=(dim, ))*jnp.sqrt(10)
covX = genCovMat(subkey2, dim)
X = jnp.dot(random.normal(subkey3, shape=(N,dim)), jnp.linalg.cholesky(covX))
p_array = batch_logistic(theta_true, X)
keys = random.split(key, N)
y_data = batch_benoulli(keys, p_array).astype(jnp.int32)
return theta_true, X, y_data
def build_grad_log_post(X, y_data, N):
"""
Builds grad_log_post
"""
@jit
def loglikelihood(theta, x_val, y_val):
return -logsumexp(jnp.array([0., (1.-2.*y_val)*jnp.dot(theta, x_val)]))
@jit
def log_prior(theta):
return -(0.5/10)*jnp.dot(theta,theta)
batch_loglik = jit(vmap(loglikelihood, in_axes=(None, 0,0)))
def log_post(theta):
return log_prior(theta) + N*jnp.mean(batch_loglik(theta, X, y_data), axis=0)
grad_log_post = jit(grad(log_post))
return grad_log_post
def build_value_and_grad_log_post(X, y_data, N):
"""
Builds grad_log_post
"""
@jit
def loglikelihood(theta, x_val, y_val):
return -logsumexp(jnp.array([0., (1.-2.*y_val)*jnp.dot(theta, x_val)]))
@jit
def log_prior(theta):
return -(0.5/10)*jnp.dot(theta,theta)
batch_loglik = jit(vmap(loglikelihood, in_axes=(None, 0,0)))
def log_post(theta):
return log_prior(theta) + N*jnp.mean(batch_loglik(theta, X, y_data), axis=0)
val_and_grad_log_post = jit(value_and_grad(log_post))
return val_and_grad_log_post
def build_batch_grad_log_post(X, y_data, N):
"""
Builds grad_log_post that takes in minibatches X and y_data
"""
@jit
def loglikelihood(theta, x_val, y_val):
return -logsumexp(jnp.array([0., (1.-2.*y_val)*jnp.dot(theta, x_val)]))
@jit
def log_prior(theta):
return -(0.5/10)*jnp.dot(theta,theta)
batch_loglik = jit(vmap(loglikelihood, in_axes=(None, 0,0)))
def log_post(theta, X, y_data):
return log_prior(theta) + N*jnp.mean(batch_loglik(theta, X, y_data), axis=0)
grad_log_post = jit(grad(log_post))
return grad_log_post