(self, act_fn)
| 109 | |
| 110 | @parameterized.parameters(["nn.gelu", "quick_gelu"]) |
| 111 | def test_clip_textual_encoder(self, act_fn): |
| 112 | act_fn_key = act_fn.replace(".", "_") |
| 113 | golden = _load_golden_jax(f"test_clip_textual_encoder_{act_fn_key}") |
| 114 | |
| 115 | model_dim = 32 |
| 116 | ff_dim = 64 |
| 117 | num_layers = 3 |
| 118 | num_heads = 8 |
| 119 | max_seq_len = 12 |
| 120 | |
| 121 | kwargs = { |
| 122 | "pad_token_id": OUR_PAD_TOKEN_ID, |
| 123 | "max_seq_len": max_seq_len, |
| 124 | "vocab_size": OUR_VOCAB_SIZE, |
| 125 | "num_layers": num_layers, |
| 126 | "model_dim": model_dim, |
| 127 | "num_heads": num_heads, |
| 128 | "feed_forward_dim": ff_dim, |
| 129 | "feed_forward_act": act_fn, |
| 130 | "dropout_rate": 0, |
| 131 | "projection_dim": model_dim, |
| 132 | } |
| 133 | |
| 134 | layer_cfg = set_text_encoder_config(**kwargs) |
| 135 | layer_cfg.set(name="test") |
| 136 | layer = layer_cfg.instantiate(parent=None) |
| 137 | |
| 138 | # Golden params cover the text_encoder subtree; initialize full params and overlay. |
| 139 | layer_params = layer.initialize_parameters_recursively(prng_key=jax.random.PRNGKey(0)) |
| 140 | layer_params.update(golden["params"]) |
| 141 | # Convert HF-format input_ids (PAD==EOS==49407) to AXLearn format (PAD=49408). |
| 142 | input_ids = _hf_to_axlearn_input_ids(golden["inputs"]["input_ids"]) |
| 143 | |
| 144 | layer_outputs, _ = F( |
| 145 | layer, |
| 146 | inputs=dict(input_batch={"text": jnp.expand_dims(jnp.asarray(input_ids), 1)}), |
| 147 | state=layer_params, |
| 148 | is_training=True, |
| 149 | prng_key=jax.random.PRNGKey(0), |
| 150 | ) |
| 151 | assert_allclose( |
| 152 | jnp.squeeze(layer_outputs[TEXT_EMBEDDINGS]), |
| 153 | jnp.asarray(golden["outputs"]["pooler_output"]), |
| 154 | ) |
| 155 | |
| 156 | |
| 157 | class TestCLIPModel(parameterized.TestCase): |
nothing calls this directly
no test coverage detected