-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathnumpy_rbf.py
More file actions
64 lines (48 loc) · 2.08 KB
/
Copy pathnumpy_rbf.py
File metadata and controls
64 lines (48 loc) · 2.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
import numpy as np
# RBF Layer
class RBF(object):
"""
Transforms incoming data using a given radial basis function:
u_{i} = rbf(||x - c_{i}|| / s_{i})
Arguments:
in_features: size of each input sample
out_features: size of each output sample
Shape:
- Input: (N, in_features) where N is an arbitrary batch size
- Output: (N, out_features) where N is an arbitrary batch size
Attributes:
centers: the learnable centres of shape (out_features, in_features).
The values are initialised from a standard normal distribution.
Normalising inputs to have mean 0 and standard deviation 1 is
recommended.
widths: the learnable scaling factors of shape (out_features).
The values are initialised as ones.
basis_func: the radial basis function used to transform the scaled
distances.
"""
def __init__(self, in_features, out_features, basis_func):
self.in_features = in_features
self.out_features = out_features
self.centers, self.widths = self.reset_parameters()
self.basis_func = basis_func
def reset_parameters(self):
centers = np.random.randn(self.out_features, self.in_features)
widths = np.ones([self.out_features, ])
return centers, widths
def eval_basis(self, input): # (B, x)
num_batch = np.shape(input)[0]
size = (num_batch, self.out_features, self.in_features) # (B, y, x)
x = np.repeat(input.reshape([-1, 1, self.in_features]), self.out_features, axis=1) # (B, y, x)
c = np.repeat(self.centers.reshape([-1, self.out_features, self.in_features]), num_batch, axis=0) # (B, y, x)
distances = np.sum((x - c) ** 2, axis=2) ** (0.5) * self.widths.reshape(-1, self.out_features) # (B, y)
return self.basis_func(distances)
# RBFs
def gaussian(alpha):
phi = np.exp(-1 * alpha ** 2)
return phi
def basis_func_dict():
"""
A helper function that returns a dictionary containing each RBF
"""
bases = {'gaussian': gaussian}
return bases