MCPcopy Create free account
hub / github.com/Tencent/NeuralNLP-NeuralClassifier / ClassificationTrainer

Class ClassificationTrainer

train.py:87–185  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

85
86
87class ClassificationTrainer(object):
88 def __init__(self, label_map, logger, evaluator, conf, loss_fn):
89 self.label_map = label_map
90 self.logger = logger
91 self.evaluator = evaluator
92 self.conf = conf
93 self.loss_fn = loss_fn
94 if self.conf.task_info.hierarchical:
95 self.hierar_relations = get_hierar_relations(
96 self.conf.task_info.hierar_taxonomy, label_map)
97
98 def train(self, data_loader, model, optimizer, stage, epoch):
99 model.update_lr(optimizer, epoch)
100 model.train()
101 return self.run(data_loader, model, optimizer, stage, epoch,
102 ModeType.TRAIN)
103
104 def eval(self, data_loader, model, optimizer, stage, epoch):
105 model.eval()
106 return self.run(data_loader, model, optimizer, stage, epoch)
107
108 def run(self, data_loader, model, optimizer, stage,
109 epoch, mode=ModeType.EVAL):
110 is_multi = False
111 # multi-label classifcation
112 if self.conf.task_info.label_type == ClassificationType.MULTI_LABEL:
113 is_multi = True
114 predict_probs = []
115 standard_labels = []
116 num_batch = data_loader.__len__()
117 total_loss = 0.
118 for batch in data_loader:
119 # hierarchical classification using hierarchy penalty loss
120 if self.conf.task_info.hierarchical:
121 logits = model(batch)
122 linear_paras = model.linear.weight
123 is_hierar = True
124 used_argvs = (self.conf.task_info.hierar_penalty, linear_paras, self.hierar_relations)
125 loss = self.loss_fn(
126 logits,
127 batch[ClassificationDataset.DOC_LABEL].to(self.conf.device),
128 is_hierar,
129 is_multi,
130 *used_argvs)
131 # hierarchical classification with HMCN
132 elif self.conf.model_name == "HMCN":
133 (global_logits, local_logits, logits) = model(batch)
134 loss = self.loss_fn(
135 global_logits,
136 batch[ClassificationDataset.DOC_LABEL].to(self.conf.device),
137 False,
138 is_multi)
139 loss += self.loss_fn(
140 local_logits,
141 batch[ClassificationDataset.DOC_LABEL].to(self.conf.device),
142 False,
143 is_multi)
144 # flat classificaiton

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected