Skip to content

Commit 026ca3f

Browse files
Arm backend: Fix Gemma3n and CLIP export tests (pytorch#22830)
Index Gemma3n audio encoder outputs to handle both tuples and ModelOutput. Run CLIP eagerly during test setup to install hidden-state hooks before strict export encounters their installation lock. cc @digantdesai @freddan80 @per @zingo @oscarandersson8218 @mansnils @Sebastian-Larsson @robell @rascani Signed-off-by: Sangwon Ha <sangwon.ha@arm.com>
1 parent e2323a0 commit 026ca3f

2 files changed

Lines changed: 22 additions & 14 deletions

File tree

‎backends/arm/test/models/Gemma3n/test_gemma3nModel.py‎

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -548,8 +548,7 @@ def __init__(self, config) -> None:
548548
self.encoder = Gemma3nAudioEncoder(config)
549549

550550
def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
551-
encodings, _ = self.encoder(x, mask)
552-
return encodings
551+
return self.encoder(x, mask)[0]
553552

554553
@staticmethod
555554
def _prepare_inputs(

‎backends/arm/test/models/stable_diffusion_3_5_large/test_CLIPTextModelWithProjection_sd35_large.py‎

Lines changed: 21 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -67,14 +67,19 @@ def create_dummy_inputs(
6767
),
6868
)
6969

70-
def create_model(
70+
def prepare_model_and_inputs(
7171
self,
7272
config,
73-
) -> SD3CLIPTextEncoderWrapper:
74-
"""Instantiate wrapped CLIPTextModelWithProjection for tests."""
75-
return SD3CLIPTextEncoderWrapper(
73+
) -> tuple[SD3CLIPTextEncoderWrapper, input_t]:
74+
"""Prepare wrapped CLIPTextModelWithProjection for export."""
75+
model = SD3CLIPTextEncoderWrapper(
7676
CLIPTextModelWithProjection(config).to(dtype=config.dtype) # type: ignore[call-arg]
7777
).eval()
78+
inputs = self.create_dummy_inputs(config)
79+
# Install Transformers' lazy hidden-state hooks before strict export.
80+
with torch.no_grad():
81+
model(*inputs)
82+
return model, inputs
7883

7984
@staticmethod
8085
def ops_after_partitioner_INT(config) -> dict[str, int]:
@@ -107,11 +112,12 @@ def test_clip_text_model_with_projection_tosa_FP(config_factory, atol):
107112
"""Run the CLIPTextModelWithProjection TOSA FP test for a given config."""
108113
test_helper = TestCLIPTextModelWithProjection()
109114
config = config_factory()
115+
model, inputs = test_helper.prepare_model_and_inputs(config)
110116

111117
with torch.no_grad():
112118
pipeline = TosaPipelineFP[input_t](
113-
test_helper.create_model(config),
114-
test_helper.create_dummy_inputs(config),
119+
model,
120+
inputs,
115121
aten_op=[],
116122
exir_op=[],
117123
use_to_edge_transform_and_lower=True,
@@ -136,11 +142,12 @@ def test_clip_text_model_with_projection_tosa_INT(config_factory, atol):
136142
"""Run the CLIPTextModelWithProjection TOSA INT test for a given config."""
137143
test_helper = TestCLIPTextModelWithProjection()
138144
config = config_factory()
145+
model, inputs = test_helper.prepare_model_and_inputs(config)
139146

140147
with torch.no_grad():
141148
pipeline = TosaPipelineINT[input_t](
142-
test_helper.create_model(config),
143-
test_helper.create_dummy_inputs(config),
149+
model,
150+
inputs,
144151
aten_op=[],
145152
exir_op=[],
146153
use_to_edge_transform_and_lower=True,
@@ -168,11 +175,12 @@ def test_clip_text_model_with_projection_vgf_no_quant(config_factory):
168175
"""Run the CLIPTextModelWithProjection VGF no-quant test."""
169176
test_helper = TestCLIPTextModelWithProjection()
170177
config = config_factory()
178+
model, inputs = test_helper.prepare_model_and_inputs(config)
171179

172180
with torch.no_grad():
173181
pipeline = VgfPipeline[input_t](
174-
test_helper.create_model(config),
175-
test_helper.create_dummy_inputs(config),
182+
model,
183+
inputs,
176184
aten_op=[],
177185
exir_op=[],
178186
use_to_edge_transform_and_lower=True,
@@ -200,11 +208,12 @@ def test_clip_text_model_with_projection_vgf_quant(config_factory, atol):
200208
"""Run the CLIPTextModelWithProjection VGF quant test."""
201209
test_helper = TestCLIPTextModelWithProjection()
202210
config = config_factory()
211+
model, inputs = test_helper.prepare_model_and_inputs(config)
203212

204213
with torch.no_grad():
205214
pipeline = VgfPipeline[input_t](
206-
test_helper.create_model(config),
207-
test_helper.create_dummy_inputs(config),
215+
model,
216+
inputs,
208217
aten_op=[],
209218
exir_op=[],
210219
use_to_edge_transform_and_lower=True,

0 commit comments

Comments
 (0)