mirror of
https://github.com/modelscope/modelscope.git
synced 2026-09-01 19:49:03 +02:00
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:
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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])
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user