mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
added debug arguments for regularization
This commit is contained in:
+3
-3
@@ -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
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user