File size: 3,970 Bytes
bf928ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
from tf_keras.layers import *
from tf_keras.models import *
from tf_keras import backend as K
from tf_keras.layers import Layer
from tf_keras import initializers


MAX_LEN_cdr = 24
NB_WORDS_cdr = 8001  # kmer1 -> 21 , kmer2 -> 401, kmer3 -> 8001
NB_cdr_number_ids = 3

MAX_LEN_ag = 2371
NB_WORDS_ag = 21  # kmer1 -> 21 , kmer2 -> 401, kmer3 -> 8001

EMBEDDING_DIM = 100

filters = 256

cdr_kernel_size = 6
cdr_pool_size = cdr_strides = 4

ag_kernel_size = 60
ag_pool_size = ag_strides = 20

lstm_size = 50

att_size = 50

dt_ratio = 0.5

class AttLayer(Layer):
    def __init__(self, attention_dim):
        self.init = initializers.RandomNormal(seed=10)
        self.supports_masking = True
        self.attention_dim = attention_dim
        super(AttLayer, self).__init__()

    def build(self, input_shape):
        assert len(input_shape) == 3

        self.W = self.add_weight(
            name="W",
            shape=(input_shape[-1], self.attention_dim),
            initializer=self.init,
            trainable=True,
        )
        self.b = self.add_weight(
            name="b",
            shape=(self.attention_dim,),
            initializer=self.init,
            trainable=True,
        )
        self.u = self.add_weight(
            name="u",
            shape=(self.attention_dim, 1),
            initializer=self.init,
            trainable=True,
        )

        super(AttLayer, self).build(input_shape)

    def compute_mask(self, inputs, mask=None):
        return mask

    def call(self, x, mask=None):
        # size of x :[batch_size, sel_len, attention_dim]
        # size of u :[batch_size, attention_dim]
        # uit = tanh(xW+b)
        uit = K.tanh(K.bias_add(K.dot(x, self.W), self.b))
        ait = K.dot(uit, self.u)
        ait = K.squeeze(ait, -1)

        ait = K.exp(ait)

        if mask is not None:
            # Cast the mask to floatX to avoid float64 upcasting in theano
            ait *= K.cast(mask, K.floatx())
        ait /= K.cast(K.sum(ait, axis=1, keepdims=True) + K.epsilon(), K.floatx())
        ait = K.expand_dims(ait)
        weighted_input = x * ait
        output = K.sum(weighted_input, axis=1)

        return output

    def compute_output_shape(self, input_shape):
        return (input_shape[0], input_shape[-1])


def get_model():
    cdrs_ids = Input(shape=(MAX_LEN_cdr,))
    cdrs_number_ids = Input(shape=(MAX_LEN_cdr,))

    ags_ids = Input(shape=(MAX_LEN_ag,))

    emb_cdr_ids = Embedding(NB_WORDS_cdr, EMBEDDING_DIM, trainable=True)(cdrs_ids)
    emb_cdr_number_ids = Embedding(NB_cdr_number_ids, EMBEDDING_DIM, trainable=True)(cdrs_number_ids)
    emb_cdr = Add()([emb_cdr_ids, emb_cdr_number_ids ])
    emb_cdr_bn = BatchNormalization()(emb_cdr)
    emb_cdr_dt = Dropout(dt_ratio)(emb_cdr_bn)

    emb_ag_ids = Embedding(NB_WORDS_ag, EMBEDDING_DIM, trainable=True)(ags_ids)
    emb_ag_bn = BatchNormalization()(emb_ag_ids)
    emb_ag_dt = Dropout(dt_ratio)(emb_ag_bn)

    cdr_conv_layer = Conv1D(filters = filters, kernel_size = cdr_kernel_size,padding = "valid",activation='relu')(emb_cdr_dt)
    cdr_max_pool_layer = MaxPooling1D(pool_size = cdr_pool_size, strides = cdr_strides)(cdr_conv_layer)

    ag_conv_layer = Conv1D(filters = filters, kernel_size = ag_kernel_size,padding = "valid",activation='relu')(emb_ag_dt)
    ag_max_pool_layer = MaxPooling1D(pool_size = ag_pool_size, strides = ag_strides)(ag_conv_layer)

    merge_layer=Concatenate(axis=1)([cdr_max_pool_layer, ag_max_pool_layer])
    bn=BatchNormalization()(merge_layer)
    dt=Dropout(dt_ratio)(bn)

    l_lstm = Bidirectional(LSTM(lstm_size, return_sequences=True))(dt)
    l_att = AttLayer(att_size)(l_lstm)

    preds = Dense(1, activation='sigmoid')(l_att)

    model = Model(inputs=[cdrs_ids, cdrs_number_ids, ags_ids],outputs= [preds])

    model.compile(loss='binary_crossentropy',optimizer='adam')

    return model