MCPcopy Create free account
hub / github.com/dmlc/parameter_server / updateWeight

Method updateWeight

src/app/linear_method/darlin.h:206–246  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

204 }
205
206 void updateWeight(int grp, SizeR range, SArray<Real> G, SArray<Real> U) {
207 CHECK_EQ(G.size(), range.size());
208 CHECK_EQ(U.size(), range.size());
209
210 Real eta = conf_.learning_rate().alpha();
211 Real lambda = conf_.penalty().lambda(0);
212 Real delta_max = bcd_conf_.GetExtension(delta_max_value);
213 auto& value = model_.value(grp);
214 auto& active_set = active_set_[grp];
215 auto& delta = delta_[grp];
216 for (size_t i = 0; i < range.size(); ++i) {
217 size_t k = i + range.begin();
218 if (!active_set.test(k)) continue;
219 Real g = G[i], u = U[i] / eta + 1e-10;
220 Real g_pos = g + lambda, g_neg = g - lambda;
221 Real& w = value[k];
222 Real d = - w, vio = 0;
223
224 if (w == 0) {
225 if (g_pos < 0) {
226 vio = - g_pos;
227 } else if (g_neg > 0) {
228 vio = g_neg;
229 } else if (g_pos > kkt_filter_threshold_ && g_neg < - kkt_filter_threshold_) {
230 active_set.clear(k);
231 kkt_filter_.mark(&w);
232 continue;
233 }
234 }
235 violation_ = std::max(violation_, vio);
236
237 if (g_pos <= u * w) {
238 d = - g_pos / u;
239 } else if (g_neg >= u * w) {
240 d = - g_neg / u;
241 }
242 d = std::min(delta[k], std::max(-delta[k], d));
243 delta[k] = newDelta(delta_max, d);
244 w += d;
245 }
246 }
247
248 virtual void evaluate(BCDProgress* prog) {
249 size_t nnz_w = 0;

Callers

nothing calls this directly

Calls 8

lambdaMethod · 0.80
testMethod · 0.80
markMethod · 0.80
sizeMethod · 0.45
alphaMethod · 0.45
valueMethod · 0.45
beginMethod · 0.45
clearMethod · 0.45

Tested by

no test coverage detected