| 16 | |
| 17 | |
| 18 | class BaseFM(BaseEstimator): |
| 19 | def __init__( |
| 20 | self, |
| 21 | n_components=10, |
| 22 | max_iter=100, |
| 23 | init_stdev=0.1, |
| 24 | learning_rate=0.01, |
| 25 | reg_v=0.1, |
| 26 | reg_w=0.5, |
| 27 | reg_w0=0.0, |
| 28 | ): |
| 29 | """Simplified factorization machines implementation using SGD optimizer.""" |
| 30 | self.reg_w0 = reg_w0 |
| 31 | self.reg_w = reg_w |
| 32 | self.reg_v = reg_v |
| 33 | self.n_components = n_components |
| 34 | self.lr = learning_rate |
| 35 | self.init_stdev = init_stdev |
| 36 | self.max_iter = max_iter |
| 37 | self.loss = None |
| 38 | self.loss_grad = None |
| 39 | |
| 40 | def fit(self, X, y=None): |
| 41 | self._setup_input(X, y) |
| 42 | # bias |
| 43 | self.wo = 0.0 |
| 44 | # Feature weights |
| 45 | self.w = np.zeros(self.n_features) |
| 46 | # Factor weights |
| 47 | self.v = np.random.normal( |
| 48 | scale=self.init_stdev, size=(self.n_features, self.n_components) |
| 49 | ) |
| 50 | self._train() |
| 51 | |
| 52 | def _train(self): |
| 53 | for epoch in range(self.max_iter): |
| 54 | y_pred = self._predict(self.X) |
| 55 | loss = self.loss_grad(self.y, y_pred) |
| 56 | w_grad = np.dot(loss, self.X) / float(self.n_samples) |
| 57 | self.wo -= self.lr * (loss.mean() + 2 * self.reg_w0 * self.wo) |
| 58 | self.w -= self.lr * w_grad + (2 * self.reg_w * self.w) |
| 59 | self._factor_step(loss) |
| 60 | |
| 61 | def _factor_step(self, loss): |
| 62 | for ix, x in enumerate(self.X): |
| 63 | for i in range(self.n_features): |
| 64 | v_grad = loss[ix] * (x.dot(self.v).dot(x[i])[0] - self.v[i] * x[i] ** 2) |
| 65 | self.v[i] -= self.lr * v_grad + (2 * self.reg_v * self.v[i]) |
| 66 | |
| 67 | def _predict(self, X=None): |
| 68 | linear_output = np.dot(X, self.w) |
| 69 | factors_output = ( |
| 70 | np.sum(np.dot(X, self.v) ** 2 - np.dot(X**2, self.v**2), axis=1) / 2.0 |
| 71 | ) |
| 72 | return self.wo + linear_output + factors_output |
| 73 | |
| 74 | |
| 75 | class FMRegressor(BaseFM): |
nothing calls this directly
no outgoing calls
no test coverage detected