refactor faq model and add MGIMN model

FAQ模型代码重构+新增FAQ MGIMN模型 

Link: https://code.alibaba-inc.com/Ali-MaaS/MaaS-lib/codereview/11595371
This commit is contained in:
tanfan.zjh
2023-02-09 08:29:19 +00:00
committed by wenmeng.zwm
parent ce4199a783
commit bb174351b3
7 changed files with 477 additions and 63 deletions

View File

@@ -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):

View File

@@ -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

View File

@@ -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

View File

@@ -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])

View File

@@ -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

View File

@@ -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)

View File

@@ -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()