Views
No views yet
1# -*- coding: utf-8 -*-
2import tensorflow as tf
3from tensorflow.keras import layers, Model
4from tensorflow.keras.optimizers import Adam
5from tensorflow.keras.losses import SparseCategoricalCrossentropy
6from transformers import (BertTokenizer, TFBertModel,
7 RobertaTokenizer, TFRobertaModel,
8 AlbertTokenizer, TFAlbertModel,
9 DebertaTokenizer, TFDebertaModel,
10 FunnelTokenizer, TFFunnelModel)
11
12
13class Transformer_EBD_Reg:
14 def __init__(self):
15 self.num_classes = 3
16 self.shared_fc1 = layers.Dense(768, activation='tanh') # Assuming 768 as the dimension of output embeddings
17 self.shared_fc2 = layers.Dense(1, activation='sigmoid')
18 self.shared_pooling = layers.GlobalMaxPool1D()
19 self.shared_output_layer = layers.Dense(self.num_classes, activation="softmax")
20 self.model_function = {'bert':self.load_bert,'albert':self.load_albert,'roberta':self.load_roberta,'deberta':self.load_deberta,'funnel_tf':self.load_funnel_tf}
21
22 def _build_model(self, model, input_shapes, weight_path):
23 inputs = [layers.Input(shape=shape, dtype=tf.int32, name=name)
24 for shape, name in zip(input_shapes.values(), input_shapes.keys())]
25 bert_output = model(*inputs).last_hidden_state
26
27 bert_output_transformed = self.shared_fc2(self.shared_fc1(bert_output))
28 bert_output_multiplied = bert_output_transformed * bert_output
29
30 norm = tf.norm(bert_output_multiplied, axis=1, keepdims=True)
31 bert_output_multiplied_normalized = bert_output_multiplied / norm
32
33 bert_output_pooled = self.shared_pooling(bert_output_multiplied_normalized)
34 output = self.shared_output_layer(bert_output_pooled)
35
36 new_model = Model(inputs=inputs, outputs=[output])
37 new_model.load_weights(weight_path)
38
39 loss = SparseCategoricalCrossentropy(from_logits=False)
40 optimizer = Adam(learning_rate=1e-5)
41 new_model.compile(optimizer=optimizer, loss=loss, metrics=['accuracy'])
42
43 return new_model
44
45 def load_bert(self,model):
46 tokenizer = BertTokenizer.from_pretrained("bert-base-uncased")
47 bert_model = TFBertModel.from_pretrained('bert-base-uncased')
48 input_shapes = {"input_ids": (None,), "attention_mask": (None,), "token_type_ids": (None,)}
49 weight_path = '{}_weights.h5'.format(model)
50 return self._build_model(bert_model, input_shapes, weight_path), tokenizer
51
52 def load_roberta(self,model):
53 tokenizer = RobertaTokenizer.from_pretrained("roberta-base")
54 roberta_model = TFRobertaModel.from_pretrained('roberta-base')
55 input_shapes = {"input_ids": (None,), "attention_mask": (None,)}
56 weight_path = '{}_weights.h5'.format(model)
57 return self._build_model(roberta_model, input_shapes, weight_path), tokenizer
58
59 def load_albert(self,model):
60 tokenizer = AlbertTokenizer.from_pretrained("albert-base-v2")
61 albert_model = TFAlbertModel.from_pretrained('albert-base-v2')
62 input_shapes = {"input_ids": (None,), "attention_mask": (None,), "token_type_ids": (None,)}
63 weight_path = '{}_weights.h5'.format(model)
64 return self._build_model(albert_model, input_shapes, weight_path), tokenizer
65
66 def load_deberta(self,model):
67 tokenizer = DebertaTokenizer.from_pretrained("microsoft/deberta-base")
68 deberta_model = TFDebertaModel.from_pretrained('microsoft/deberta-base')
69 input_shapes = {"input_ids": (None,), "attention_mask": (None,), "token_type_ids": (None,)}
70 weight_path = '{}_weights.h5'.format(model)
71 return self._build_model(deberta_model, input_shapes, weight_path), tokenizer
72
73 def load_funnel_tf(self,model):
74 tokenizer = FunnelTokenizer.from_pretrained('funnel-transformer/small')
75 funnel_model = TFFunnelModel.from_pretrained('funnel-transformer/small')
76 input_shapes = {"input_ids": (None,), "attention_mask": (None,), "token_type_ids": (None,)}
77 weight_path = '{}_weights.h5'.format(model)
78 return self._build_model(funnel_model, input_shapes, weight_path), tokenizer
79
80 def load_weights(self, model):
81 return self.model_function[model](model)
82
83
84
85if __name__ == '__main__':
86 transformer_model = Transformer_EBD_Reg()
87 roberta_model, roberta_tokenizer = transformer_model.load_weights('roberta')
88 roberta_model.summary()