MCPcopy Create free account
hub / github.com/apple/axlearn / test_clip_visual_encoder

Method test_clip_visual_encoder

axlearn/vision/clip_test.py:63–108  ·  view source on GitHub ↗
(self, act_fn)

Source from the content-addressed store, hash-verified

61
62 @parameterized.parameters(["nn.gelu", "quick_gelu"])
63 def test_clip_visual_encoder(self, act_fn):
64 act_fn_key = act_fn.replace(".", "_")
65 golden = _load_golden_jax(f"test_clip_visual_encoder_{act_fn_key}")
66
67 model_dim = 32
68 ff_dim = 64
69 image_size = 16
70 patch_size = 4
71 num_layers = 3
72 num_heads = 8
73
74 kwargs = {
75 "num_layers": num_layers,
76 "model_dim": model_dim,
77 "num_heads": num_heads,
78 "feed_forward_dim": ff_dim,
79 "feed_forward_act": act_fn,
80 "image_size": (image_size, image_size),
81 "patch_size": (patch_size, patch_size),
82 "dropout_rate": 0,
83 "projection_dim": model_dim,
84 }
85 layer_cfg = set_vision_encoder_config(**kwargs)
86 layer_cfg.set(name="test")
87 layer = layer_cfg.instantiate(parent=None)
88
89 # Golden params cover the image_encoder subtree; initialize full params and overlay.
90 layer_params = layer.initialize_parameters_recursively(prng_key=jax.random.PRNGKey(0))
91 layer_params.update(golden["params"])
92 pixel_values = golden["inputs"]["pixel_values"]
93
94 layer_outputs, _ = F(
95 layer,
96 inputs=dict(
97 input_batch=dict(
98 image=jnp.expand_dims(jnp.einsum("bchw->bhwc", jnp.asarray(pixel_values)), 1)
99 )
100 ),
101 state=layer_params,
102 is_training=True,
103 prng_key=jax.random.PRNGKey(0),
104 )
105 assert_allclose(
106 jnp.squeeze(layer_outputs["pooled_features"]),
107 jnp.asarray(golden["outputs"]["pooler_output"]),
108 )
109
110 @parameterized.parameters(["nn.gelu", "quick_gelu"])
111 def test_clip_textual_encoder(self, act_fn):

Callers

nothing calls this directly

Calls 8

assert_allcloseFunction · 0.90
_load_golden_jaxFunction · 0.85
replaceMethod · 0.80
setMethod · 0.45
instantiateMethod · 0.45
updateMethod · 0.45

Tested by

no test coverage detected