From bb174351b362b7600bce860ac045ce95e94680be Mon Sep 17 00:00:00 2001 From: "tanfan.zjh" Date: Thu, 9 Feb 2023 08:29:19 +0000 Subject: [PATCH] refactor faq model and add MGIMN model MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit FAQ模型代码重构+新增FAQ MGIMN模型 Link: https://code.alibaba-inc.com/Ali-MaaS/MaaS-lib/codereview/11595371 --- modelscope/metrics/accuracy_metric.py | 3 + .../nlp/structbert/faq_question_answering.py | 452 +++++++++++++++--- .../nlp/faq_question_answering_pipeline.py | 2 + .../faq_question_answering_preprocessor.py | 21 + .../nlp/faq_question_answering_trainer.py | 18 +- .../pipelines/test_faq_question_answering.py | 9 + .../test_finetune_faq_question_answering.py | 35 +- 7 files changed, 477 insertions(+), 63 deletions(-) diff --git a/modelscope/metrics/accuracy_metric.py b/modelscope/metrics/accuracy_metric.py index b1976d8e..2327a9c7 100644 --- a/modelscope/metrics/accuracy_metric.py +++ b/modelscope/metrics/accuracy_metric.py @@ -8,6 +8,7 @@ from modelscope.metainfo import Metrics from modelscope.outputs import OutputKeys from modelscope.utils.chinese_utils import remove_space_between_chinese_chars from modelscope.utils.registry import default_group +from modelscope.utils.tensor_utils import torch_nested_numpify from .base import Metric from .builder import METRICS, MetricKeys @@ -36,8 +37,10 @@ class AccuracyMetric(Metric): eval_results = outputs[key] break assert type(ground_truths) == type(eval_results) + ground_truths = torch_nested_numpify(ground_truths) for truth in ground_truths: self.labels.append(truth) + eval_results = torch_nested_numpify(eval_results) for result in eval_results: if isinstance(truth, str): if isinstance(result, list): diff --git a/modelscope/models/nlp/structbert/faq_question_answering.py b/modelscope/models/nlp/structbert/faq_question_answering.py index c5cd3061..bedd09e8 100644 --- a/modelscope/models/nlp/structbert/faq_question_answering.py +++ b/modelscope/models/nlp/structbert/faq_question_answering.py @@ -9,6 +9,7 @@ import torch import torch.nn as nn import torch.nn.functional as F from torch import Tensor +from torch.nn import BCEWithLogitsLoss from modelscope.metainfo import Models from modelscope.models.builder import MODELS @@ -17,6 +18,9 @@ from modelscope.models.nlp.task_models.task_model import BaseTaskModel from modelscope.outputs import FaqQuestionAnsweringOutput from modelscope.utils.config import Config, ConfigFields from modelscope.utils.constant import ModelFile, Tasks +from modelscope.utils.logger import get_logger + +logger = get_logger() activations = { 'relu': F.relu, @@ -88,9 +92,6 @@ class MetricsLayer(nn.Module): return self.args.metrics def forward(self, query, protos): - """ query : [bsz, n_query, dim] - support : [bsz, n_query, n_cls, dim] | [bsz, n_cls, dim] - """ if self.args.metrics == 'cosine': supervised_dists = self.cosine_similarity(query, protos) if self.training: @@ -102,8 +103,6 @@ class MetricsLayer(nn.Module): return supervised_dists def cosine_similarity(self, x, y): - # x=[bsz, n_query, dim] - # y=[bsz, n_cls, dim] n_query = x.shape[0] n_cls = y.shape[0] dim = x.shape[-1] @@ -155,42 +154,68 @@ class PoolingLayer(nn.Module): return self.pooling(x, mask) +class Alignment(nn.Module): + + def __init__(self): + super().__init__() + + def _attention(self, a, b): + return torch.matmul(a, b.transpose(1, 2)) + + def forward(self, a, b, mask_a, mask_b): + attn = self._attention(a, b) + mask = torch.matmul(mask_a.float(), mask_b.transpose(1, 2).float()) + mask = mask.bool() + attn.masked_fill_(~mask, -1e4) + return attn + + +def _create_args(model_config, hidden_size): + metric = model_config.get('metric', 'cosine') + pooling_method = model_config.get('pooling', 'avg') + Arg = namedtuple( + 'args', + ['metrics', 'proj_hidden_size', 'hidden_size', 'dropout', 'pooling']) + args = Arg( + metrics=metric, + proj_hidden_size=hidden_size, + hidden_size=hidden_size, + dropout=0.0, + pooling=pooling_method) + return args + + @MODELS.register_module( Tasks.faq_question_answering, module_name=Models.structbert) class SbertForFaqQuestionAnswering(BaseTaskModel): _backbone_prefix = '' + PROTO_NET = 'protonet' + MGIMN_NET = 'mgimnnet' @classmethod def _instantiate(cls, **kwargs): - model = cls(kwargs.get('model_dir')) - model.load_checkpoint(kwargs.get('model_dir')) + model_dir = kwargs.pop('model_dir') + model = cls(model_dir, **kwargs) + model.load_checkpoint(model_dir) return model def __init__(self, model_dir, *args, **kwargs): super().__init__(model_dir, *args, **kwargs) - backbone_cfg = SbertConfig.from_pretrained(model_dir) - self.bert = SbertModel(backbone_cfg) - model_config = Config.from_file( os.path.join(model_dir, ModelFile.CONFIGURATION)).get(ConfigFields.model, {}) + model_config.update(kwargs) - metric = model_config.get('metric', 'cosine') - pooling_method = model_config.get('pooling', 'avg') - - Arg = namedtuple('args', [ - 'metrics', 'proj_hidden_size', 'hidden_size', 'dropout', 'pooling' - ]) - args = Arg( - metrics=metric, - proj_hidden_size=self.bert.config.hidden_size, - hidden_size=self.bert.config.hidden_size, - dropout=0.0, - pooling=pooling_method) - - self.metrics_layer = MetricsLayer(args) - self.pooling = PoolingLayer(args) + network_name = model_config.get('network', self.PROTO_NET) + if network_name == self.PROTO_NET: + network = ProtoNet(backbone_cfg, model_config) + elif network_name == self.MGIMN_NET: + network = MGIMNNet(backbone_cfg, model_config) + else: + raise NotImplementedError(network_name) + logger.info(f'faq task build {network_name} network') + self.network = network def forward(self, input: Dict[str, Tensor]) -> FaqQuestionAnsweringOutput: """ @@ -242,18 +267,87 @@ class SbertForFaqQuestionAnswering(BaseTaskModel): support = input['support'] query_mask = input['query_attention_mask'] support_mask = input['support_attention_mask'] - - n_query = query.shape[0] - n_support = support.shape[0] - support_labels = input['support_labels'] + logits, scores = self.network(query, support, query_mask, support_mask, + support_labels) + + if 'labels' in input: + query_labels = input['labels'] + num_cls = torch.max(support_labels) + 1 + loss = self._compute_loss(logits, query_labels, num_cls) + pred_labels = torch.argmax(scores, dim=1) + return FaqQuestionAnsweringOutput( + loss=loss, logits=scores, labels=pred_labels).to_dict() + else: + return FaqQuestionAnsweringOutput(scores=scores) + + def _compute_loss(self, logits, target, num_cls): + onehot_labels = get_onehot_labels(target, num_cls) + loss = BCEWithLogitsLoss(reduction='mean')(logits, onehot_labels) + return loss + + def forward_sentence_embedding(self, inputs): + return self.network.sentence_embedding(inputs) + + def load_checkpoint(self, + model_local_dir, + default_dtype=None, + load_state_fn=None, + **kwargs): + ckpt_file = os.path.join(model_local_dir, 'pytorch_model.bin') + state_dict = torch.load(ckpt_file, map_location='cpu') + # compatible with the old checkpoints + new_state_dict = {} + for var_name, var_value in state_dict.items(): + new_var_name = var_name + if not str(var_name).startswith('network'): + new_var_name = f'network.{var_name}' + new_state_dict[new_var_name] = var_value + if default_dtype is not None: + torch.set_default_dtype(default_dtype) + + missing_keys, unexpected_keys, mismatched_keys, error_msgs = self._load_checkpoint( + new_state_dict, + load_state_fn=load_state_fn, + ignore_mismatched_sizes=True, + _fast_init=True, + ) + + return { + 'missing_keys': missing_keys, + 'unexpected_keys': unexpected_keys, + 'mismatched_keys': mismatched_keys, + 'error_msgs': error_msgs, + } + + +def get_onehot_labels(target, num_cls): + target = target.view(-1, 1) + size = target.shape[0] + target_oh = torch.zeros(size, num_cls).to(target) + target_oh.scatter_(dim=1, index=target, value=1) + return target_oh.view(size, num_cls).float() + + +class ProtoNet(nn.Module): + + def __init__(self, backbone_config, model_config): + super(ProtoNet, self).__init__() + self.bert = SbertModel(backbone_config) + args = _create_args(model_config, self.bert.config.hidden_size) + self.metrics_layer = MetricsLayer(args) + self.pooling = PoolingLayer(args) + + def __call__(self, query, support, query_mask, support_mask, + support_labels): + n_query = query.shape[0] + num_cls = torch.max(support_labels) + 1 - onehot_labels = self._get_onehot_labels(support_labels, n_support, - num_cls) + onehot_labels = get_onehot_labels(support_labels, num_cls) input_ids = torch.cat([query, support]) input_mask = torch.cat([query_mask, support_mask], dim=0) - pooled_representation = self.forward_sentence_embedding({ + pooled_representation = self.sentence_embedding({ 'input_ids': input_ids, 'attention_mask': @@ -269,29 +363,9 @@ class SbertForFaqQuestionAnswering(BaseTaskModel): scores = torch.sigmoid(logits) else: scores = logits - if 'labels' in input: - query_labels = input['labels'] - loss = self._compute_loss(logits, query_labels, num_cls) - _, pred_labels = torch.max(scores, dim=1) - return FaqQuestionAnsweringOutput( - loss=loss, logits=scores).to_dict() - else: - return FaqQuestionAnsweringOutput(scores=scores) + return logits, scores - def _compute_loss(self, logits, target, num_cls): - from torch.nn import CrossEntropyLoss - logits = logits.view([-1, num_cls]) - target = target.reshape(-1) - loss = CrossEntropyLoss(reduction='mean')(logits, target) - return loss - - def _get_onehot_labels(self, labels, support_size, num_cls): - labels_ = labels.view(support_size, 1) - target_oh = torch.zeros(support_size, num_cls).to(labels) - target_oh.scatter_(dim=1, index=labels_, value=1) - return target_oh.view(support_size, num_cls).float() - - def forward_sentence_embedding(self, inputs: Dict[str, Tensor]): + def sentence_embedding(self, inputs: Dict[str, Tensor]): input_ids = inputs['input_ids'] input_mask = inputs['attention_mask'] if not isinstance(input_ids, Tensor): @@ -304,3 +378,273 @@ class SbertForFaqQuestionAnswering(BaseTaskModel): input_mask = input_mask.unsqueeze(-1) pooled_representation = self.pooling(last_hidden_states, input_mask) return pooled_representation + + +class MGIMNNet(nn.Module): + # default use class_level_interaction only + INSTANCE_LEVEL_INTERACTION = 'instance_level_interaction' + EPISODE_LEVEL_INTERACTION = 'episode_level_interaction' + + def __init__(self, backbone_config, model_config): + super(MGIMNNet, self).__init__() + self.bert = SbertModel(backbone_config) + self.model_config = model_config + self.alignment = Alignment() + hidden_size = self.bert.config.hidden_size + use_instance_level_interaction = self.safe_get( + self.INSTANCE_LEVEL_INTERACTION, True) + use_episode_level_interaction = self.safe_get( + self.EPISODE_LEVEL_INTERACTION, True) + output_size = 1 + int(use_instance_level_interaction) + int( + use_episode_level_interaction) + logger.info( + f'faq MGIMN model class-level-interaction:true, instance-level-interaction:{use_instance_level_interaction}, \ + episode-level-interaction:{use_episode_level_interaction}') + self.fuse_proj = LinearProjection( + hidden_size + hidden_size * 3 * output_size, + hidden_size, + activation='relu') + args = _create_args(model_config, hidden_size) + self.pooling = PoolingLayer(args) + new_args = args._replace(pooling='avg') + self.avg_pooling = PoolingLayer(new_args) + + self.instance_compare_layer = torch.nn.Sequential( + LinearProjection(hidden_size * 4, hidden_size, activation='relu')) + + self.prediction = torch.nn.Sequential( + LinearProjection(hidden_size * 2, hidden_size, activation='relu'), + nn.Dropout(0), LinearProjection(hidden_size, 1)) + + def __call__(self, query, support, query_mask, support_mask, + support_labels): + z_query, z_support = self.context_embedding(query, support, query_mask, + support_mask) + n_cls = int(torch.max(support_labels)) + 1 + n_query, sent_len = query.shape + n_support = support.shape[0] + k_shot = n_support // n_cls + + q_params, s_params = { + 'n_cls': n_cls, + 'k_shot': k_shot + }, { + 'n_query': n_query + } + if self.safe_get(self.INSTANCE_LEVEL_INTERACTION, True): + ins_z_query, ins_z_support = self._instance_level_interaction( + z_query, query_mask, z_support, support_mask) + q_params['ins_z_query'] = ins_z_query + s_params['ins_z_support'] = ins_z_support + + cls_z_query, cls_z_support = self._class_level_interaction( + z_query, query_mask, z_support, support_mask, n_cls) + q_params['cls_z_query'] = cls_z_query + s_params['cls_z_support'] = cls_z_support + if self.safe_get(self.EPISODE_LEVEL_INTERACTION, True): + eps_z_query, eps_z_support = self._episode_level_interaction( + z_query, query_mask, z_support, support_mask) + q_params['eps_z_query'] = eps_z_query + s_params['eps_z_support'] = eps_z_support + fused_z_query = self._fuse_query(z_query, **q_params) + fused_z_support = self._fuse_support(z_support, **s_params) + query_mask_expanded = query_mask.unsqueeze(1).repeat( + 1, n_support, 1).view(n_query * n_support, sent_len, 1) + support_mask_expanded = support_mask.unsqueeze(0).repeat( + n_query, 1, 1).view(n_query * n_support, sent_len, 1) + Q = self.pooling(fused_z_query, query_mask_expanded) + S = self.pooling(fused_z_support, support_mask_expanded) + matching_feature = self._instance_compare(Q, S, n_query, n_cls, k_shot) + logits = self.prediction(matching_feature) + logits = logits.view(n_query, n_cls) + return logits, torch.sigmoid(logits) + + def _instance_compare(self, Q, S, n_query, n_cls, k_shot): + z_dim = Q.shape[-1] + S = S.view(n_query, n_cls * k_shot, z_dim) + Q = Q.view(n_query, k_shot * n_cls, z_dim) + cat_features = torch.cat([Q, S, Q * S, (Q - S).abs()], dim=-1) + instance_matching_feature = self.instance_compare_layer(cat_features) + instance_matching_feature = instance_matching_feature.view( + n_query, n_cls, k_shot, z_dim) + cls_matching_feature_mean = instance_matching_feature.mean(2) + cls_matching_feature_max, _ = instance_matching_feature.max(2) + cls_matching_feature = torch.cat( + [cls_matching_feature_mean, cls_matching_feature_max], dim=-1) + return cls_matching_feature + + def _instance_level_interaction(self, z_query, query_mask, z_support, + support_mask): + n_query, sent_len, z_dim = z_query.shape + n_support = z_support.shape[0] + z_query = z_query.unsqueeze(1).repeat(1, n_support, 1, + 1).view(n_query * n_support, + sent_len, z_dim) + query_mask = query_mask.unsqueeze(1).repeat(1, n_support, 1).view( + n_query * n_support, sent_len, 1) + z_support = z_support.unsqueeze(0).repeat(n_query, 1, 1, 1).view( + n_query * n_support, sent_len, z_dim) + support_mask = support_mask.unsqueeze(0).repeat(n_query, 1, 1, 1).view( + n_query * n_support, sent_len, 1) + attn = self.alignment(z_query, z_support, query_mask, support_mask) + attn_a = F.softmax(attn, dim=1) + attn_b = F.softmax(attn, dim=2) + ins_support = torch.matmul(attn_a.transpose(1, 2), z_query) + ins_query = torch.matmul(attn_b, z_support) + return ins_query, ins_support + + def _class_level_interaction(self, z_query, query_mask, z_support, + support_mask, n_cls): + z_support_ori = z_support + support_mask_ori = support_mask + + n_query, sent_len, z_dim = z_query.shape + n_support = z_support.shape[0] + k_shot = n_support // n_cls + + # class-based query encoding + z_query = z_query.unsqueeze(1).repeat(1, n_cls, 1, + 1).view(n_query * n_cls, + sent_len, z_dim) + query_mask = query_mask.unsqueeze(1).unsqueeze(-1).repeat( + 1, n_cls, 1, 1).view(n_query * n_cls, sent_len, 1) + z_support = z_support.unsqueeze(0).repeat(n_query, 1, 1, 1).view( + n_query * n_cls, k_shot * sent_len, z_dim) + support_mask = support_mask.unsqueeze(0).unsqueeze(-1).repeat( + n_query, 1, 1, 1).view(n_query * n_cls, k_shot * sent_len, 1) + attn = self.alignment(z_query, z_support, query_mask, support_mask) + attn_b = F.softmax(attn, dim=2) + cls_query = torch.matmul(attn_b, z_support) + cls_query = cls_query.view(n_query, n_cls, sent_len, z_dim) + + # class-based support encoding + z_support = z_support_ori.view(n_cls, k_shot * sent_len, z_dim) + support_mask = support_mask_ori.view(n_cls, k_shot * sent_len, 1) + attn = self.alignment(z_support, z_support, support_mask, support_mask) + attn_b = F.softmax(attn, dim=2) + cls_support = torch.matmul(attn_b, z_support) + cls_support = cls_support.view(n_cls * k_shot, sent_len, z_dim) + return cls_query, cls_support + + def _episode_level_interaction(self, z_query, query_mask, z_support, + support_mask): + z_support_ori = z_support + support_mask_ori = support_mask + + n_query, sent_len, z_dim = z_query.shape + n_support = z_support.shape[0] + + # episode-based query encoding + query_mask = query_mask.view(n_query, sent_len, 1) + z_support = z_support.unsqueeze(0).repeat(n_query, 1, 1, 1).view( + n_query, n_support * sent_len, z_dim) + support_mask = support_mask.unsqueeze(0).unsqueeze(-1).repeat( + n_query, 1, 1, 1).view(n_query, n_support * sent_len, 1) + attn = self.alignment(z_query, z_support, query_mask, support_mask) + attn_b = F.softmax(attn, dim=2) + eps_query = torch.matmul(attn_b, z_support) + + # episode-based support encoding + z_support2 = z_support_ori.view(1, n_support * sent_len, + z_dim).repeat(n_support, 1, 1) + support_mask = support_mask_ori.view(1, n_support * sent_len, + 1).repeat(n_support, 1, 1) + attn = self.alignment(z_support_ori, z_support2, + support_mask_ori.unsqueeze(-1), support_mask) + attn_b = F.softmax(attn, dim=2) + eps_support = torch.matmul(attn_b, z_support2) + eps_support = eps_support.view(n_support, sent_len, z_dim) + return eps_query, eps_support + + def _fuse_query(self, + x, + n_cls, + k_shot, + ins_z_query=None, + cls_z_query=None, + eps_z_query=None): + n_query, sent_len, z_dim = x.shape + assert cls_z_query is not None + cls_features = cls_z_query.unsqueeze(2).repeat( + 1, 1, k_shot, 1, 1).view(n_cls * k_shot * n_query, sent_len, z_dim) + x = x.unsqueeze(1).repeat(1, n_cls * k_shot, 1, + 1).view(n_cls * k_shot * n_query, sent_len, + z_dim) + features = [ + x, cls_features, x * cls_features, (x - cls_features).abs() + ] + if ins_z_query is not None: + features.extend( + [ins_z_query, ins_z_query * x, (ins_z_query - x).abs()]) + if eps_z_query is not None: + eps_z_query = eps_z_query.unsqueeze(1).repeat( + 1, n_cls * k_shot, 1, 1).view(n_cls * k_shot * n_query, + sent_len, z_dim) + features.extend( + [eps_z_query, eps_z_query * x, (eps_z_query - x).abs()]) + features = torch.cat(features, dim=-1) + fusion_feat = self.fuse_proj(features) + return fusion_feat + + def _fuse_support(self, + x, + n_query, + ins_z_support=None, + cls_z_support=None, + eps_z_support=None): + assert cls_z_support is not None + n_support, sent_len, z_dim = x.shape + x = x.unsqueeze(0).repeat(n_query, 1, 1, + 1).view(n_support * n_query, sent_len, z_dim) + cls_features = cls_z_support.unsqueeze(0).repeat( + n_query, 1, 1, 1).view(n_support * n_query, sent_len, z_dim) + features = [ + x, cls_features, x * cls_features, (x - cls_features).abs() + ] + if ins_z_support is not None: + features.extend( + [ins_z_support, ins_z_support * x, (ins_z_support - x).abs()]) + if eps_z_support is not None: + eps_z_support = eps_z_support.unsqueeze(0).repeat( + n_query, 1, 1, 1).view(n_query * n_support, sent_len, z_dim) + features.extend( + [eps_z_support, eps_z_support * x, (eps_z_support - x).abs()]) + features = torch.cat(features, dim=-1) + fusion_feat = self.fuse_proj(features) + return fusion_feat + + def context_embedding(self, query, support, query_mask, support_mask): + n_query = query.shape[0] + n_support = support.shape[0] + x = torch.cat([query, support], dim=0) + x_mask = torch.cat([query_mask, support_mask], dim=0) + last_hidden_state = self.bert(x, x_mask).last_hidden_state + z_dim = last_hidden_state.shape[-1] + sent_len = last_hidden_state.shape[-2] + z_query = last_hidden_state[:n_query].view([n_query, sent_len, z_dim]) + z_support = last_hidden_state[n_query:].view( + [n_support, sent_len, z_dim]) + return z_query, z_support + + def sentence_embedding(self, inputs: Dict[str, Tensor]): + input_ids = inputs['input_ids'] + input_mask = inputs['attention_mask'] + if not isinstance(input_ids, Tensor): + input_ids = torch.IntTensor(input_ids) + if not isinstance(input_mask, Tensor): + input_mask = torch.IntTensor(input_mask) + rst = self.bert(input_ids, input_mask) + last_hidden_states = rst.last_hidden_state + if len(input_mask.shape) == 2: + input_mask = input_mask.unsqueeze(-1) + pooled_representation = self.avg_pooling(last_hidden_states, + input_mask) + return pooled_representation + + def safe_get(self, k, default=None): + try: + return self.model_config.get(k, default) + except Exception as e: + logger.debug(f'{k} not in model_config, use default:{default}') + logger.debug(e) + return default diff --git a/modelscope/pipelines/nlp/faq_question_answering_pipeline.py b/modelscope/pipelines/nlp/faq_question_answering_pipeline.py index 4e9ebf01..0b2ba199 100644 --- a/modelscope/pipelines/nlp/faq_question_answering_pipeline.py +++ b/modelscope/pipelines/nlp/faq_question_answering_pipeline.py @@ -43,6 +43,8 @@ class FaqQuestionAnsweringPipeline(Pipeline): if preprocessor is None: self.preprocessor = Preprocessor.from_pretrained( self.model.model_dir, **kwargs) + if hasattr(self.model, 'eval'): + self.model.eval() def _sanitize_parameters(self, **pipeline_parameters): return pipeline_parameters, pipeline_parameters, pipeline_parameters diff --git a/modelscope/preprocessors/nlp/faq_question_answering_preprocessor.py b/modelscope/preprocessors/nlp/faq_question_answering_preprocessor.py index eb8c501f..157d204d 100644 --- a/modelscope/preprocessors/nlp/faq_question_answering_preprocessor.py +++ b/modelscope/preprocessors/nlp/faq_question_answering_preprocessor.py @@ -57,6 +57,8 @@ class FaqQuestionAnsweringTransformersPreprocessor(Preprocessor): self.support_set = support_set self.label_in_support_set = label_in_support_set self.text_in_support_set = text_in_support_set + # support non-prototype network + self.pad_support = kwargs.get('pad_support', False) def pad(self, samples, max_len): result = [] @@ -79,6 +81,23 @@ class FaqQuestionAnsweringTransformersPreprocessor(Preprocessor): ] + self.tokenizer.convert_tokens_to_ids( self.tokenizer.tokenize(text)) + [self.tokenizer.sep_token_id] + def preprocess(self, support_set): + label_to_samples = {} + for item in support_set: + label = item[self.label_in_support_set] + if label not in label_to_samples: + label_to_samples[label] = [] + label_to_samples[label].append(item) + max_cnt = 0 + for label, samples in label_to_samples.items(): + if len(samples) > max_cnt: + max_cnt = len(samples) + new_support_set = [] + for label, samples in label_to_samples.items(): + new_support_set.extend( + samples + [samples[0] for _ in range(max_cnt - len(samples))]) + return new_support_set + @type_assert(object, Dict) def __call__(self, data: Dict[str, Any], **preprocessor_param) -> Dict[str, Any]: @@ -93,6 +112,8 @@ class FaqQuestionAnsweringTransformersPreprocessor(Preprocessor): if not isinstance(queryset, list): queryset = [queryset] supportset = data[self.support_set] + if self.pad_support: + supportset = self.preprocess(supportset) supportset = sorted( supportset, key=lambda d: d[self.label_in_support_set]) diff --git a/modelscope/trainers/nlp/faq_question_answering_trainer.py b/modelscope/trainers/nlp/faq_question_answering_trainer.py index a4a78cf7..dc6f0426 100644 --- a/modelscope/trainers/nlp/faq_question_answering_trainer.py +++ b/modelscope/trainers/nlp/faq_question_answering_trainer.py @@ -64,6 +64,9 @@ class EpisodeSampler(torch.utils.data.BatchSampler): self.episode = n_iter domain_label_sampleid = {} bad_sample_ids = self.get_bad_sampleids(dataset) + if dataset.mode == 'train': + logger.info( + f'num. of bad sample ids:{len(bad_sample_ids)}/{len(dataset)}') for sample_index, sample in enumerate(dataset): if sample_index in bad_sample_ids: continue @@ -95,7 +98,9 @@ class EpisodeSampler(torch.utils.data.BatchSampler): data_size += len(tokens) if dataset.mode == 'train': logger.info( - f'{dataset.mode}: label size:{total}, data size:{data_size}') + f'{dataset.mode}: label size:{total}, data size:{data_size}, \ + domain_size:{len(self.domain_label_tokens)}') + self.mode = dataset.mode def __iter__(self): for i in range(self.episode): @@ -109,18 +114,21 @@ class EpisodeSampler(torch.utils.data.BatchSampler): list(self.domain_label_tokens[domain].keys())) N = min(self.n_way, len(all_labels)) labels = np.random.choice( - all_labels, size=min(N, len(all_labels)), replace=False) + all_labels, size=min(N, len(all_labels)), + replace=False).tolist() batch = [] for label in labels[:N]: candidates = self.domain_label_tokens[domain][label] - K = min(len(candidates), int((self.k_shot + self.r_query))) - tmp = np.random.choice(candidates, size=K, replace=False) + num_samples = self.k_shot + self.r_query + K = min(len(candidates), int(num_samples)) + tmp = np.random.choice( + candidates, size=K, replace=False).tolist() batch.extend(tmp) batch = [int(n) for n in batch] yield batch def _get_field(self, obj, key, default=None): - value = getattr(obj, key, default) or obj.get(key, default) + value = obj.get(key, default) if value is not None: return str(value) return None diff --git a/tests/pipelines/test_faq_question_answering.py b/tests/pipelines/test_faq_question_answering.py index 5aca5a1a..dc29e385 100644 --- a/tests/pipelines/test_faq_question_answering.py +++ b/tests/pipelines/test_faq_question_answering.py @@ -21,6 +21,7 @@ class FaqQuestionAnsweringTest(unittest.TestCase, DemoCompatibilityCheck): def setUp(self) -> None: self.task = Tasks.faq_question_answering self.model_id = 'damo/nlp_structbert_faq-question-answering_chinese-base' + self.mgimn_model_id = 'damo/nlp_mgimn_faq-question-answering_chinese-base' self.model_id_multilingual = 'damo/nlp_faq-question-answering_multilingual-base' param = { @@ -90,6 +91,14 @@ class FaqQuestionAnsweringTest(unittest.TestCase, DemoCompatibilityCheck): pipeline_ins = pipeline(task=Tasks.faq_question_answering) print(pipeline_ins(self.param, max_seq_length=20)) + @unittest.skipUnless(test_level() >= 2, 'skip test in current test level') + def test_run_with_mgimn_model(self): + pipeline_ins = pipeline( + task=Tasks.faq_question_answering, + model=self.mgimn_model_id, + model_revision='v1.0.0') + print(pipeline_ins(self.param, max_seq_length=20)) + @unittest.skipUnless(test_level() >= 2, 'skip test in current test level') def test_sentence_embedding(self): pipeline_ins = pipeline(task=Tasks.faq_question_answering) diff --git a/tests/trainers/test_finetune_faq_question_answering.py b/tests/trainers/test_finetune_faq_question_answering.py index 01c34b63..54b0a840 100644 --- a/tests/trainers/test_finetune_faq_question_answering.py +++ b/tests/trainers/test_finetune_faq_question_answering.py @@ -32,6 +32,7 @@ class TestFinetuneFaqQuestionAnswering(unittest.TestCase): }] } model_id = 'damo/nlp_structbert_faq-question-answering_chinese-base' + mgimn_model_id = 'damo/nlp_mgimn_faq-question-answering_chinese-base' def setUp(self): print(('Testing %s.%s' % (type(self).__name__, self._testMethodName))) @@ -43,7 +44,7 @@ class TestFinetuneFaqQuestionAnswering(unittest.TestCase): shutil.rmtree(self.tmp_dir) super().tearDown() - def build_trainer(self): + def build_trainer(self, model_id, revision): train_dataset = MsDataset.load( 'jd', namespace='DAMO_NLP', split='train').remap_columns({'sentence': 'text'}) @@ -51,7 +52,7 @@ class TestFinetuneFaqQuestionAnswering(unittest.TestCase): 'jd', namespace='DAMO_NLP', split='validation').remap_columns({'sentence': 'text'}) - cfg: Config = read_config(self.model_id, revision='v1.0.1') + cfg: Config = read_config(model_id, revision) cfg.train.train_iters_per_epoch = 50 cfg.evaluation.val_iters_per_epoch = 2 cfg.train.seed = 1234 @@ -75,7 +76,7 @@ class TestFinetuneFaqQuestionAnswering(unittest.TestCase): trainer = build_trainer( Trainers.faq_question_answering_trainer, default_args=dict( - model=self.model_id, + model=model_id, work_dir=self.tmp_dir, train_dataset=train_dataset, eval_dataset=eval_dataset, @@ -84,7 +85,7 @@ class TestFinetuneFaqQuestionAnswering(unittest.TestCase): @unittest.skipUnless(test_level() >= 0, 'skip test in current test level') def test_faq_model_finetune(self): - trainer = self.build_trainer() + trainer = self.build_trainer(self.model_id, 'v1.0.1') trainer.train() evaluate_result = trainer.evaluate() self.assertAlmostEqual(evaluate_result['accuracy'], 0.95, delta=0.1) @@ -106,6 +107,32 @@ class TestFinetuneFaqQuestionAnswering(unittest.TestCase): self.assertAlmostEqual( result_after['output'][0][0]['score'], 0.8, delta=0.2) + @unittest.skipUnless(test_level() >= 0, 'skip test in current test level') + def test_faq_mgimn_model_finetune(self): + trainer = self.build_trainer(self.mgimn_model_id, 'v1.0.0') + trainer.train() + evaluate_result = trainer.evaluate() + self.assertAlmostEqual(evaluate_result['accuracy'], 0.75, delta=0.1) + + results_files = os.listdir(self.tmp_dir) + self.assertIn(ModelFile.TRAIN_OUTPUT_DIR, results_files) + + output_dir = os.path.join(self.tmp_dir, ModelFile.TRAIN_OUTPUT_DIR) + pipeline_ins = pipeline( + task=Tasks.faq_question_answering, + model=self.mgimn_model_id, + model_revision='v1.0.0') + result_before = pipeline_ins(self.param) + self.assertEqual(result_before['output'][0][0]['label'], '1') + self.assertAlmostEqual( + result_before['output'][0][0]['score'], 0.9, delta=0.2) + pipeline_ins = pipeline( + task=Tasks.faq_question_answering, model=output_dir) + result_after = pipeline_ins(self.param) + self.assertEqual(result_after['output'][0][0]['label'], '1') + self.assertAlmostEqual( + result_after['output'][0][0]['score'], 0.9, delta=0.2) + if __name__ == '__main__': unittest.main()