| 409 | |
| 410 | |
| 411 | class 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 |
no outgoing calls
no test coverage detected
searching dependent graphs…