added debug arguments for regularization

This commit is contained in:
Gaurav Baruah
2017-03-27 14:42:35 -04:00
parent 1bc571b635
commit 84910639a2
2 changed files with 12 additions and 10 deletions
+3 -3
View File
@@ -85,9 +85,9 @@ if __name__ == "__main__":
# debugging arguments
ap.add_argument('--debugSingleBatch', action="store_true", help="will stop program after training 1 input batch")
ap.add_argument('--no_ext_feats', action="store_true", help="will not include external features in the model")
ap.add_argument('--num_conv_filters', help="the number of convolution channels (lesser is faster)", default=100, type=int)
ap.add_argument('--no_ext_feats', action="store_true", help="will not include external features in the model")
ap.add_argument('--no_loss_reg', help="no loss regularization", action="store_true")
args = ap.parse_args()
@@ -106,7 +106,7 @@ if __name__ == "__main__":
torch.set_num_threads(args.num_threads)
trainer = Trainer(net, args.eta, args.mom)
trainer = Trainer(net, args.eta, args.mom, args.no_loss_reg)
best_accuracy = 0.0
best_model = 0
+9 -7
View File
@@ -27,11 +27,12 @@ logger.addHandler(ch)
class Trainer(object):
def __init__(self, model, eta, mom):
def __init__(self, model, eta, mom, no_loss_reg):
self.reg = 1e-5
self.no_loss_reg = no_loss_reg
self.model = model
self.criterion = nn.CrossEntropyLoss()
self.optimizer = optim.SGD(self.model.parameters(), lr=eta, momentum=mom, weight_decay=self.reg)
self.optimizer = optim.SGD(self.model.parameters(), lr=eta, momentum=mom, weight_decay=(0 if no_loss_reg else self.reg) )
def regularize_loss(self, loss):
@@ -60,8 +61,9 @@ class Trainer(object):
logger.debug('loss after criterion {}'.format(loss))
# NOTE: regularizing location 1
# loss = self.regularize_loss(loss)
# logger.debug('loss after regularizing {}'.format(loss))
# if not self.no_loss_reg:
# loss = self.regularize_loss(loss)
# logger.debug('loss after regularizing {}'.format(loss))
loss.backward()
@@ -69,10 +71,10 @@ class Trainer(object):
#logger.debug('params {}'.format([p for p in self.model.parameters()]))
logger.debug('params grads {}'.format([p.grad for p in self.model.parameters()]))
# NOTE: regularizing location 2. It would seem that location 1 is correct?
loss = self.regularize_loss(loss)
logger.debug('loss after regularizing {}'.format(loss))
if not self.no_loss_reg:
loss = self.regularize_loss(loss)
logger.debug('loss after regularizing {}'.format(loss))
self.optimizer.step()