-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathpreprocess.py
More file actions
26 lines (20 loc) · 723 Bytes
/
Copy pathpreprocess.py
File metadata and controls
26 lines (20 loc) · 723 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
import numpy as np
"""This file contains some functions related to preprocessing."""
def onehot(output_labels):
return np.squeeze((np.unique(output_labels) == output_labels[:, None]).astype(np.float32))
def get_fans(shape):
return np.prod(shape[:-1]), shape[-1]
def xavier_init(shape):
fan_in, fan_out = get_fans(shape)
dev = np.sqrt(6.0 / (fan_in + fan_out))
return np.random.uniform(-dev, dev, shape)
def xavier_init_normal(shape):
fan_in, fan_out = get_fans(shape)
dev = np.sqrt(3.0 / (fan_in + fan_out))
return np.random.normal(size=shape, scale=dev)
def randomize(data, labels):
order = np.arange(data.shape[0])
np.random.shuffle(order)
data = data[order]
labels = labels[order]
return data, labels