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

Method train

rat-sql-gap/seq2struct/commands/train.py:102–205  ·  view source on GitHub ↗
(self, config, modeldir)

Source from the content-addressed store, hash-verified

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(
134 {"model": self.model, "optimizer": optimizer}, keep_every_n=self.train_config.keep_every_n)
135 last_step = saver.restore(modeldir, map_location=self.device)
136
137 if "pretrain" in config and last_step == 0:
138 pretrain_config = config["pretrain"]
139 _path = pretrain_config["pretrained_path"]
140 _step = pretrain_config["checkpoint_step"]
141 pretrain_step = saver.restore(_path, step=_step, map_location=self.device, item_keys=["model"])
142 saver.save(modeldir, pretrain_step) # for evaluating pretrained models
143 last_step = pretrain_step
144
145 # 3. Get training data somewhere
146 with self.data_random:
147 train_data = self.model_preproc.dataset('train')
148 train_data_loader = self._yield_batches_from_epochs(
149 torch.utils.data.DataLoader(
150 train_data,
151 batch_size=self.train_config.batch_size,
152 shuffle=True,
153 drop_last=True,
154 collate_fn=lambda x: x))
155 train_eval_data_loader = torch.utils.data.DataLoader(
156 train_data,
157 batch_size=self.train_config.eval_batch_size,
158 collate_fn=lambda x: x)
159

Callers 2

mainFunction · 0.95
_eval_modelMethod · 0.45

Calls 10

restoreMethod · 0.95
saveMethod · 0.95
_eval_modelMethod · 0.95
appendMethod · 0.80
logMethod · 0.80
datasetMethod · 0.45
compute_lossMethod · 0.45
stepMethod · 0.45
update_lrMethod · 0.45

Tested by

no test coverage detected