(self, net, param_init_net, param_info, fp32_update=False)
| 500 | self.weight_decay = weight_decay |
| 501 | |
| 502 | def _run(self, net, param_init_net, param_info, fp32_update=False): |
| 503 | |
| 504 | fp32_update_flag = 0 |
| 505 | param_name = str(param_info.blob) |
| 506 | |
| 507 | # should only be triggered in FP16 training by SpatialBN, which |
| 508 | # requires FP32 params in CuDNN. |
| 509 | if param_name.find("spatbn") != -1: |
| 510 | fp32_update = True |
| 511 | |
| 512 | if fp32_update: |
| 513 | # doing a 32bit update |
| 514 | # Have to assume param_info.blob is FP32 as there is no way |
| 515 | # (that i currently know of) to query a blob's type in python |
| 516 | fp32_update_flag = 1 |
| 517 | param = param_info.blob |
| 518 | param_fp32 = param_info.blob |
| 519 | else: |
| 520 | if param_info.blob_copy is None: |
| 521 | # doing a 32bit update |
| 522 | # Have to assume param_info.blob is FP32 as there is no way |
| 523 | # (that i currently know of) to query a blob's type in python |
| 524 | fp32_update_flag = 1 |
| 525 | param = param_info.blob |
| 526 | param_fp32 = param_info.blob |
| 527 | else: |
| 528 | if core.DataType.FLOAT in param_info.blob_copy: |
| 529 | param = param_info.blob |
| 530 | param_fp32 = param_info.blob_copy[core.DataType.FLOAT] |
| 531 | elif core.DataType.FLOAT16 in param_info.blob_copy: |
| 532 | param = param_info.blob_copy[core.DataType.FLOAT16] |
| 533 | param_fp32 = param_info.blob |
| 534 | else: |
| 535 | AssertionError( |
| 536 | "Unrecognized parameter format to be updated " |
| 537 | "by FP16 Optimizer. Parameter: {}".format(param_info.name) |
| 538 | ) |
| 539 | |
| 540 | grad = param_info.grad |
| 541 | |
| 542 | if self.base_learning_rate == 0: |
| 543 | return |
| 544 | assert ( |
| 545 | self.base_learning_rate > 0 |
| 546 | ), "Expect positive base learning rate, got {}".format(self.base_learning_rate) |
| 547 | |
| 548 | lr, _ = self.build_lr( |
| 549 | net, |
| 550 | param_init_net, |
| 551 | base_learning_rate=-self.base_learning_rate, |
| 552 | policy=self.policy, |
| 553 | **(self.init_kwargs) |
| 554 | ) |
| 555 | |
| 556 | momentum_data_fp32 = param_init_net.ConstantFill( |
| 557 | param_fp32, str(param) + "_momentum_fp32", value=0.0 |
| 558 | ) |
| 559 |
nothing calls this directly
no test coverage detected