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

Class FP16SgdOptimizer

caffe2/python/optimizer.py:481–586  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

479
480
481class FP16SgdOptimizer(SgdOptimizer):
482 def __init__(
483 self,
484 base_learning_rate=0.1,
485 momentum=0.0,
486 policy="fixed",
487 nesterov=True,
488 weight_decay=0.0001,
489 sparse_dedup_aggregator=None,
490 **kwargs
491 ):
492 super().__init__(
493 base_learning_rate=base_learning_rate,
494 policy=policy,
495 momentum=momentum,
496 nesterov=nesterov,
497 sparse_dedup_aggregator=sparse_dedup_aggregator,
498 **kwargs
499 )
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 )

Callers 1

build_fp16_sgdFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…