Function
build_trainer
(model, optimizer, loss, train_loader, val_loader, \
device, flags, global_state)
Source from the content-addressed store, hash-verified
| 117 | |
| 118 | |
| 119 | def build_trainer(model, optimizer, loss, train_loader, val_loader, \ |
| 120 | device, flags, global_state): |
| 121 | if flags.Global.algorithm in ['CRNN', 'FAN', 'GRCNN', 'DAN', 'SAR', 'SATRN']: |
| 122 | trainer = TrainerRec( |
| 123 | device=device, |
| 124 | model=model, |
| 125 | optimizer=optimizer, |
| 126 | loss=loss, |
| 127 | val_loader=val_loader, |
| 128 | train_loader=train_loader, |
| 129 | flags=flags, |
| 130 | global_state=global_state |
| 131 | ) |
| 132 | return trainer |
Tested by
no test coverage detected