MCPcopy Create free account
hub / github.com/awslabs/gap-text2sql / Trainer

Class Trainer

rat-sql-gap/seq2struct/commands/train.py:76–237  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

74 self.log_file.flush()
75
76class Trainer:
77 def __init__(self, logger, config):
78 if torch.cuda.is_available():
79 self.device = torch.device('cuda')
80 else:
81 self.device = torch.device('cpu')
82
83 self.logger = logger
84 self.train_config = registry.instantiate(TrainConfig, config['train'])
85 self.data_random = random_state.RandomContext(self.train_config.data_seed)
86 self.model_random = random_state.RandomContext(self.train_config.model_seed)
87
88 self.init_random = random_state.RandomContext(self.train_config.init_seed)
89 with self.init_random:
90 # 0. Construct preprocessors
91 self.model_preproc = registry.instantiate(
92 registry.lookup('model', config['model']).Preproc,
93 config['model'],
94 unused_keys=('name',))
95 self.model_preproc.load()
96
97 # 1. Construct model
98 self.model = registry.construct('model', config['model'],
99 unused_keys=('encoder_preproc', 'decoder_preproc'), preproc=self.model_preproc, device=self.device)
100 self.model.to(self.device)
101
102 def train(self, config, modeldir):
103 # slight difference here vs. unrefactored train: The init_random starts over here. Could be fixed if it was important by saving random state at end of init
104 with self.init_random:
105 # We may be able to move optimizer and lr_scheduler to __init__ instead. Empirically it works fine. I think that's because saver.restore
106 # resets the state by calling optimizer.load_state_dict.
107 # But, if there is no saved file yet, I think this is not true, so might need to reset the optimizer manually?
108 # For now, just creating it from scratch each time is safer and appears to be the same speed, but also means you have to pass in the config to train which is kind of ugly.
109
110 # TODO: not nice
111 if config["optimizer"].get("name", None) == 'bertAdamw':
112 bert_params = list(self.model.encoder.bert_model.parameters())
113 assert len(bert_params) > 0
114 non_bert_params = []
115 for name, _param in self.model.named_parameters():
116 if "bert" not in name:
117 non_bert_params.append(_param)
118 assert len(non_bert_params) + len(bert_params) == len(list(self.model.parameters()))
119
120 optimizer = registry.construct('optimizer', config['optimizer'], non_bert_params=non_bert_params, \
121 bert_params=bert_params)
122 lr_scheduler = registry.construct( 'lr_scheduler',
123 config.get('lr_scheduler', {'name': 'noop'}),
124 param_groups=[optimizer.non_bert_param_group, \
125 optimizer.bert_param_group])
126 else:
127 optimizer = registry.construct('optimizer', config['optimizer'], params=self.model.parameters())
128 lr_scheduler = registry.construct( 'lr_scheduler',
129 config.get('lr_scheduler', {'name': 'noop'}),
130 param_groups=optimizer.param_groups)
131
132 # 2. Restore model parameters
133 saver = saver_mod.Saver(

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected