MCPcopy Create free account
hub / github.com/pytorch/pytorch / MultiPrecisionSgdOptimizer

Class MultiPrecisionSgdOptimizer

caffe2/python/optimizer.py:411–478  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

409
410
411class MultiPrecisionSgdOptimizer(SgdOptimizer):
412 def __init__(
413 self,
414 base_learning_rate=0.1,
415 momentum=0.0,
416 policy="fixed",
417 nesterov=True,
418 sparse_dedup_aggregator=None,
419 **kwargs
420 ):
421 super().__init__(
422 base_learning_rate=base_learning_rate,
423 policy=policy,
424 momentum=momentum,
425 nesterov=nesterov,
426 sparse_dedup_aggregator=sparse_dedup_aggregator,
427 **kwargs
428 )
429
430 def _run(self, net, param_init_net, param_info):
431 param = param_info.blob
432 param_fp32 = (
433 param_info.blob_copy[core.DataType.FLOAT]
434 if param_info.blob_copy is not None
435 else None
436 )
437
438 # If we have a straight fp32 parameter, run the base class
439 if param_fp32 is None:
440 return SgdOptimizer._run(self, net, param_init_net, param_info)
441
442 grad = param_info.grad
443 if self.base_learning_rate == 0:
444 return
445 assert (
446 self.base_learning_rate > 0
447 ), "Expect positive base learning rate, got {}".format(self.base_learning_rate)
448
449 lr, _ = self.build_lr(
450 net,
451 param_init_net,
452 base_learning_rate=-self.base_learning_rate,
453 policy=self.policy,
454 **(self.init_kwargs)
455 )
456
457 momentum_data = param_init_net.ConstantFill(
458 param_fp32, str(param) + "_momentum", value=0.0
459 )
460 self._aux_params.local.append(momentum_data)
461
462 assert not isinstance(
463 grad, core.GradientSlice
464 ), "MultiPrecisionSgd does not support sparse gradients"
465
466 # Copy gradient to fp32
467 grad_fp32 = net.HalfToFloat(grad, grad + "_fp32")
468

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…