mirror of
https://github.com/wassname/Castor.git
synced 2026-08-22 11:40:35 +08:00
* Refactor datasets * Update evaluators * Update trainers * Update main and MP-CNN model * Add serialization util * Fix bugs * Refactoring for NCE to use new parent class
26 lines
972 B
Python
26 lines
972 B
Python
class Evaluator(object):
|
|
"""
|
|
Evaluates a model on a Dataset, using metrics specific to the Dataset.
|
|
"""
|
|
|
|
def __init__(self, dataset_cls, model, embedding, data_loader, batch_size, device):
|
|
self.dataset_cls = dataset_cls
|
|
self.model = model
|
|
self.embedding = embedding
|
|
self.data_loader = data_loader
|
|
self.batch_size = batch_size
|
|
self.device = device
|
|
|
|
def get_sentence_embeddings(self, batch):
|
|
sent1 = self.embedding(batch.sentence_1).transpose(1, 2)
|
|
sent2 = self.embedding(batch.sentence_2).transpose(1, 2)
|
|
return sent1, sent2
|
|
|
|
def get_scores(self):
|
|
"""
|
|
Get the scores used to evaluate the model.
|
|
Should return ([score1, score2, ..], [score1_name, score2_name, ...]).
|
|
The first score is the primary score used to determine if the model has improved.
|
|
"""
|
|
raise NotImplementedError('Evaluator subclass needs to implement get_score')
|