INNER CODE UNIT · Python
TransformScheduler
timeseriesAI/tsai · tsai/callback/core.py:22
class TransformScheduler(Callback):
"A callback to schedule batch transforms during training based on a function (sched_lin, sched_exp, sched_cos (default), etc)"
def __init__(self, schedule_func:callable, show_plot:bool=False):
self.schedule_func,self.show_plot = schedule_func,show_plot
self.mult = []
def before_fit(self):
for pct in np.linspace(0, 1, len(self.dls.train) * self.n_epoch): self.mult.append(self.schedule_func(pct))
# get initial magnitude values and update initial value
self.mag = []
self.mag_tfms = []
for t in self.dls.after_batch:
if hasattr(t, 'magnitude'):
self.mag.append(t.magnitude)
t.magnitude *= self.mult[0]
self.mag_tfms.append(t)
def after_batch(self):