MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / __init__

Method __init__

SwissArmyTransformer/sat/ops/fused_ema_adam.py:65–95  ·  view source on GitHub ↗
(self,
                 params,
                 lr=1e-3,
                 bias_correction=True,
                 betas=(0.9, 0.999),
                 eps=1e-8,
                 adam_w_mode=True,
                 weight_decay=0.,
                 amsgrad=False,
                 set_grad_none=True,
                 ema_decay=0.9999,
                 use_num_upates=True
                 )

Source from the content-addressed store, hash-verified

63 """
64
65 def __init__(self,
66 params,
67 lr=1e-3,
68 bias_correction=True,
69 betas=(0.9, 0.999),
70 eps=1e-8,
71 adam_w_mode=True,
72 weight_decay=0.,
73 amsgrad=False,
74 set_grad_none=True,
75 ema_decay=0.9999,
76 use_num_upates=True
77 ):
78
79 if amsgrad:
80 raise RuntimeError('FusedAdam does not support the AMSGrad variant.')
81 defaults = dict(lr=lr, bias_correction=bias_correction, betas=betas, eps=eps, weight_decay=weight_decay)
82 super(FusedEmaAdam, self).__init__(params, defaults)
83 self.adam_w_mode = 1 if adam_w_mode else 0
84 self.set_grad_none = set_grad_none
85
86 fused_ema_adam_cuda = FusedEmaAdamBuilder().jit_load()
87 # Skip buffer
88 self._dummy_overflow_buf = get_accelerator().IntTensor([0])
89 self.multi_tensor_ema_adam = fused_ema_adam_cuda.multi_tensor_ema_adam
90 self.ema_decay = ema_decay
91 if use_num_upates:
92 self.num_updates = 0
93 else:
94 self.num_updates = -1
95 self.collected_params = []
96
97 def zero_grad(self):
98 if self.set_grad_none:

Callers

nothing calls this directly

Calls 2

FusedEmaAdamBuilderClass · 0.90
jit_loadMethod · 0.80

Tested by

no test coverage detected