CoolFace
Apppublic

RyanTietjen/Paper-Fragmentation

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
model.py71 linesDownload Raw Back to root
1"""
2Ryan Tietjen
3Sep 2024
4Create best model for the demo
5"""
6import tensorflow as tf
7from keras import layers
8import tensorflow_hub as hub
9
10class EmbeddingLayer(layers.Layer):
11    def __init__(self, **kwargs):
12        super().__init__(**kwargs)
13        # Hardcode the module URL directly within the layer
14        # self.module_url = "https://tfhub.dev/google/universal-sentence-encoder/4"
15        url = "https://tfhub.dev/google/universal-sentence-encoder/4"
16        self.embed_model = hub.KerasLayer(url, trainable=False, name="universal_sentence_encoder")
17
18    def call(self, inputs):
19        return self.embed_model(inputs)
20
21    def get_config(self):
22        config = super().get_config()
23        # The URL is now a fixed part of the layer, so it can be included in the config for completeness
24        # config.update({'module_url': self.module_url})
25        return config
26
27def create_token_model(token_embed):
28    input_layer = layers.Input(shape=[], dtype=tf.string)
29    embedding_layer = EmbeddingLayer()
30    token_embeddings = embedding_layer(input_layer)
31    output_layer = layers.Dense(128, activation="relu")(token_embeddings)
32    model = tf.keras.Model(input_layer, output_layer)
33    return model
34
35def create_character_vectorizer_model(char_embed, char_vectorizer):
36    input_layer = layers.Input(shape=(1,), dtype=tf.string)
37    char_vectors = char_vectorizer(input_layer) # vectorize text inputs
38    char_embedding = char_embed(char_vectors) # create embedding
39    output_layer = layers.Bidirectional(layers.LSTM(32))(char_embedding)
40    model = tf.keras.Model(input_layer, output_layer)
41    return model
42
43def create_line_number_model(input_shape, name):
44    input_layer = layers.Input(shape=(input_shape,), dtype=tf.int32, name=name)
45    output_layer = layers.Dense(32, activation="relu")(input_layer)
46    model = tf.keras.Model(input_layer, output_layer)
47    return model
48
49
50def tribrid_model(num_classes, token_embed, char_embed, text_vectorizer):
51    
52    token_model = create_token_model(token_embed)
53    character_vectorizer_model = create_character_vectorizer_model(char_embed, text_vectorizer)
54    line_number_model = create_line_number_model(15, "line_number")
55    total_lines_model = create_line_number_model(20, "total_lines")
56
57    hybrid_model = layers.Concatenate(name="hybrid")([token_model.output,
58                                                      character_vectorizer_model.output])
59    
60    dense_layer = layers.Dense(256, activation="relu")(hybrid_model)
61    dense_layer = layers.Dropout(0.5)(dense_layer)
62
63    tribrid_model = layers.Concatenate(name="tribrid") ([line_number_model.output, total_lines_model.output, dense_layer])
64    output_layer = layers.Dense(num_classes, activation="softmax")(tribrid_model)
65
66    model = tf.keras.Model([line_number_model.input, total_lines_model.input, token_model.input, character_vectorizer_model.input], output_layer)
67
68    model.compile(loss=tf.keras.losses.CategoricalCrossentropy(label_smoothing=0.2),
69                  optimizer=tf.keras.optimizers.Adam(),
70                  metrics=["accuracy"])
71    return model