Skip to content

Commit 4a612fa

Browse files
committed
Merge remote-tracking branch 'upstream/master' into validate-preset-assets
2 parents 3eafec3 + 41faf96 commit 4a612fa

22 files changed

Lines changed: 175 additions & 67 deletions

keras_hub/src/models/basnet/basnet_test.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -54,7 +54,6 @@ def test_saved_model(self):
5454
input_data=self.images,
5555
)
5656

57-
@pytest.mark.skip(reason="TODO: Bug with BASNet liteRT export")
5857
def test_litert_export(self):
5958
self.run_litert_export_test(
6059
cls=BASNetImageSegmenter,

keras_hub/src/models/d_fine/d_fine_object_detector_test.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -174,6 +174,10 @@ def test_saved_model(self):
174174
input_data=self.images,
175175
)
176176

177+
@pytest.mark.xfail(
178+
condition=keras.backend.backend() == "torch",
179+
reason="D-FINE's multi-scale features hit a torch.export shape guard.",
180+
)
177181
def test_litert_export(self):
178182
backbone = DFineBackbone(**self.base_backbone_kwargs)
179183
init_kwargs = {

keras_hub/src/models/deeplab_v3/deeplab_v3_segmenter_test.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -72,7 +72,8 @@ def test_saved_model(self):
7272
)
7373

7474
@pytest.mark.skip(
75-
reason="TODO: Bug with DeepLabV3ImageSegmenter liteRT export"
75+
reason="TODO: DeepLabV3ImageSegmenter LiteRT export fails with "
76+
"'symbolic tf.Tensor used as a Python bool' while tracing for export."
7677
)
7778
def test_litert_export(self):
7879
self.run_litert_export_test(

keras_hub/src/models/f_net/f_net_text_classifier_test.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import os
22

33
import pytest
4+
from keras import backend
45

56
from keras_hub.src.models.f_net.f_net_backbone import FNetBackbone
67
from keras_hub.src.models.f_net.f_net_text_classifier import FNetTextClassifier
@@ -57,6 +58,10 @@ def test_saved_model(self):
5758
input_data=self.input_data,
5859
)
5960

61+
@pytest.mark.xfail(
62+
condition=backend.backend() == "torch",
63+
reason="litert-torch has no lowering for aten.complex (from ops.fft2).",
64+
)
6065
def test_litert_export(self):
6166
# F-Net does NOT use padding_mask - it only uses token_ids and
6267
# segment_ids. Don't add padding_mask to input_data.

keras_hub/src/models/flux/flux_backbone_test.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import pytest
2+
from keras import backend
23
from keras import ops
34

45
from keras_hub.src.models.clip.clip_text_encoder import CLIPTextEncoder
@@ -84,6 +85,10 @@ def test_saved_model(self):
8485
input_data=self.input_data,
8586
)
8687

88+
@pytest.mark.xfail(
89+
condition=backend.backend() == "torch",
90+
reason="torch.export guard from Flux's dynamic num_heads reshape.",
91+
)
8792
def test_litert_export(self):
8893
self.run_litert_export_test(
8994
cls=FluxBackbone,

keras_hub/src/models/gpt_neo_x/gpt_neo_x_causal_lm_test.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -112,13 +112,9 @@ def test_saved_model(self):
112112
)
113113

114114
def test_litert_export(self):
115-
pytest.skip(reason="TODO: Fix TFLite export bug for GPTNeoX")
116115
self.run_litert_export_test(
117116
cls=GPTNeoXCausalLM,
118117
init_kwargs=self.init_kwargs,
119118
input_data=self.input_data,
120-
output_thresholds={
121-
"max": 1e-3,
122-
"mean": 1e-4,
123-
}, # More lenient thresholds for numerical differences
119+
output_thresholds={"*": {"max": 1e-3, "mean": 1e-4}},
124120
)

keras_hub/src/models/gpt_oss/gpt_oss_causal_lm_test.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from unittest.mock import patch
22

33
import pytest
4+
from keras import backend
45
from keras import ops
56

67
from keras_hub.src.models.gpt_oss.gpt_oss_backbone import GptOssBackbone
@@ -113,6 +114,10 @@ def test_saved_model(self):
113114
input_data=self.input_data,
114115
)
115116

117+
@pytest.mark.xfail(
118+
condition=backend.backend() == "torch",
119+
reason="litert-torch NHWC rewriter has no lowering for aten.amax.",
120+
)
116121
def test_litert_export(self):
117122
self.run_litert_export_test(
118123
cls=GptOssCausalLM,

keras_hub/src/models/mixtral/mixtral_causal_lm_test.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import os
22
from unittest.mock import patch
33

4+
import keras
45
import pytest
56
from keras import ops
67

@@ -107,6 +108,10 @@ def test_saved_model(self):
107108
input_data=self.input_data,
108109
)
109110

111+
@pytest.mark.xfail(
112+
condition=keras.backend.backend() == "torch",
113+
reason="litert-torch cannot lower aten._assert_async from MoE routing.",
114+
)
110115
def test_litert_export(self):
111116
self.run_litert_export_test(
112117
cls=MixtralCausalLM,

keras_hub/src/models/moonshine/moonshine_audio_to_text_test.py

Lines changed: 0 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -145,9 +145,6 @@ def test_saved_model(self):
145145
input_data=self.input_data,
146146
)
147147

148-
@pytest.mark.skip(
149-
reason="TODO: Bug with MoonshineAudioToText liteRT export"
150-
)
151148
def test_litert_export(self):
152149
self.run_litert_export_test(
153150
cls=MoonshineAudioToText,

keras_hub/src/models/qwen3_moe/qwen3_moe_causal_lm_test.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33

44
os.environ["KERAS_BACKEND"] = "jax"
55

6+
import keras
67
import pytest
78
from keras import ops
89

@@ -124,6 +125,10 @@ def test_saved_model(self):
124125
input_data=self.input_data,
125126
)
126127

128+
@pytest.mark.xfail(
129+
condition=keras.backend.backend() == "torch",
130+
reason="litert-torch cannot lower aten._assert_async from MoE routing.",
131+
)
127132
def test_litert_export(self):
128133
self.run_litert_export_test(
129134
cls=Qwen3MoeCausalLM,

0 commit comments

Comments
 (0)