(self, act_fn)
| 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): |
nothing calls this directly
no test coverage detected