Skip to content

Commit 2909d49

Browse files
Add EDRec (keras-team#2514)
* initial code dump * fix dtype error * code reformat * fix tests * code reformat * remove prints * address gemini comments * code reformat * fix api gen error * fix tests * pre-commit fix * fix torch GPU * fix torch TPU error * torch gpu
1 parent 2858857 commit 2909d49

6 files changed

Lines changed: 981 additions & 0 deletions

File tree

keras_hub/api/models/__init__.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -211,6 +211,12 @@
211211
from keras_hub.src.models.distil_bert.distil_bert_tokenizer import (
212212
DistilBertTokenizer as DistilBertTokenizer,
213213
)
214+
from keras_hub.src.models.edrec.edrec_backbone import (
215+
EdRecBackbone as EdRecBackbone,
216+
)
217+
from keras_hub.src.models.edrec.edrec_seq2seq_lm import (
218+
EdRecSeq2SeqLM as EdRecSeq2SeqLM,
219+
)
214220
from keras_hub.src.models.efficientnet.efficientnet_backbone import (
215221
EfficientNetBackbone as EfficientNetBackbone,
216222
)
Lines changed: 147 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,147 @@
1+
import keras
2+
3+
from keras_hub.src.api_export import keras_hub_export
4+
from keras_hub.src.models.backbone import Backbone
5+
from keras_hub.src.models.edrec.edrec_layers import EdRecDecoderBlock
6+
from keras_hub.src.models.edrec.edrec_layers import EdRecEncoderBlock
7+
8+
9+
@keras_hub_export("keras_hub.models.EdRecBackbone")
10+
class EdRecBackbone(Backbone):
11+
"""EdRec Backbone model.
12+
13+
Args:
14+
vocab_size: int, size of the vocabulary.
15+
num_layers_enc: int, number of encoder layers.
16+
num_layers_dec: int, number of decoder layers.
17+
hidden_dim: int, hidden dimension (d_model).
18+
intermediate_dim: int, intermediate dimension (d_ff).
19+
num_heads: int, number of attention heads.
20+
dropout: float, dropout rate.
21+
epsilon: float, epsilon for simple RMSNorm.
22+
"""
23+
24+
def __init__(
25+
self,
26+
vocab_size,
27+
num_layers_enc,
28+
num_layers_dec,
29+
hidden_dim,
30+
intermediate_dim,
31+
num_heads,
32+
dropout=0.0,
33+
epsilon=1e-6,
34+
dtype=None,
35+
**kwargs,
36+
):
37+
# === Layers ===
38+
self.embedding = keras.layers.Embedding(
39+
input_dim=vocab_size,
40+
output_dim=hidden_dim,
41+
dtype=dtype,
42+
name="embedding",
43+
)
44+
self.encoder_layers = []
45+
for i in range(num_layers_enc):
46+
self.encoder_layers.append(
47+
EdRecEncoderBlock(
48+
hidden_dim=hidden_dim,
49+
num_heads=num_heads,
50+
intermediate_dim=intermediate_dim,
51+
dropout_rate=dropout,
52+
epsilon=epsilon,
53+
dtype=dtype,
54+
name=f"encoder_layer_{i}",
55+
)
56+
)
57+
self.decoder_layers = []
58+
for i in range(num_layers_dec):
59+
self.decoder_layers.append(
60+
EdRecDecoderBlock(
61+
hidden_dim=hidden_dim,
62+
num_heads=num_heads,
63+
intermediate_dim=intermediate_dim,
64+
dropout_rate=dropout,
65+
epsilon=epsilon,
66+
dtype=dtype,
67+
name=f"decoder_layer_{i}",
68+
)
69+
)
70+
71+
# === Functional Model ===
72+
encoder_token_ids = keras.Input(
73+
shape=(None,), dtype="int32", name="encoder_token_ids"
74+
)
75+
decoder_token_ids = keras.Input(
76+
shape=(None,), dtype="int32", name="decoder_token_ids"
77+
)
78+
encoder_padding_mask = keras.Input(
79+
shape=(None,), dtype="bool", name="encoder_padding_mask"
80+
)
81+
decoder_padding_mask = keras.Input(
82+
shape=(None,), dtype="bool", name="decoder_padding_mask"
83+
)
84+
85+
# Encoder
86+
x_enc = self.embedding(encoder_token_ids)
87+
88+
for layer in self.encoder_layers:
89+
x_enc = layer(
90+
x_enc,
91+
padding_mask=encoder_padding_mask,
92+
)
93+
94+
# Decoder
95+
x_dec = self.embedding(decoder_token_ids)
96+
for layer in self.decoder_layers:
97+
x_dec, _, _ = layer(
98+
x_dec,
99+
encoder_outputs=x_enc,
100+
decoder_padding_mask=decoder_padding_mask,
101+
encoder_padding_mask=encoder_padding_mask,
102+
)
103+
104+
super().__init__(
105+
inputs={
106+
"encoder_token_ids": encoder_token_ids,
107+
"decoder_token_ids": decoder_token_ids,
108+
"encoder_padding_mask": encoder_padding_mask,
109+
"decoder_padding_mask": decoder_padding_mask,
110+
},
111+
outputs={
112+
"encoder_sequence_output": x_enc,
113+
"decoder_sequence_output": x_dec,
114+
},
115+
dtype=dtype,
116+
**kwargs,
117+
)
118+
119+
# === Config ===
120+
self.vocab_size = vocab_size
121+
self.num_layers_enc = num_layers_enc
122+
self.num_layers_dec = num_layers_dec
123+
self.hidden_dim = hidden_dim
124+
self.intermediate_dim = intermediate_dim
125+
self.num_heads = num_heads
126+
self.dropout = dropout
127+
self.epsilon = epsilon
128+
129+
def get_config(self):
130+
config = super().get_config()
131+
config.update(
132+
{
133+
"vocab_size": self.vocab_size,
134+
"num_layers_enc": self.num_layers_enc,
135+
"num_layers_dec": self.num_layers_dec,
136+
"hidden_dim": self.hidden_dim,
137+
"intermediate_dim": self.intermediate_dim,
138+
"num_heads": self.num_heads,
139+
"dropout": self.dropout,
140+
"epsilon": self.epsilon,
141+
}
142+
)
143+
return config
144+
145+
@property
146+
def token_embedding(self):
147+
return self.embedding
Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
import pytest
2+
from keras import ops
3+
4+
from keras_hub.src.models.edrec.edrec_backbone import EdRecBackbone
5+
from keras_hub.src.tests.test_case import TestCase
6+
7+
8+
class EdRecBackboneTest(TestCase):
9+
def setUp(self):
10+
self.init_kwargs = {
11+
"vocab_size": 10,
12+
"num_layers_enc": 2,
13+
"num_layers_dec": 2,
14+
"num_heads": 2,
15+
"hidden_dim": 4,
16+
"intermediate_dim": 8,
17+
"dropout": 0.0,
18+
}
19+
self.input_data = {
20+
"encoder_token_ids": ops.ones((2, 5), dtype="int32"),
21+
"encoder_padding_mask": ops.zeros((2, 5), dtype="int32"),
22+
"decoder_token_ids": ops.ones((2, 5), dtype="int32"),
23+
"decoder_padding_mask": ops.zeros((2, 5), dtype="int32"),
24+
}
25+
26+
def test_backbone_basics(self):
27+
self.run_backbone_test(
28+
cls=EdRecBackbone,
29+
init_kwargs=self.init_kwargs,
30+
input_data=self.input_data,
31+
run_mixed_precision_check=False,
32+
expected_output_shape={
33+
"encoder_sequence_output": (2, 5, 4),
34+
"decoder_sequence_output": (2, 5, 4),
35+
},
36+
)
37+
38+
@pytest.mark.large
39+
def test_saved_model(self):
40+
self.run_model_saving_test(
41+
cls=EdRecBackbone,
42+
init_kwargs=self.init_kwargs,
43+
input_data=self.input_data,
44+
)

0 commit comments

Comments
 (0)