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

Method test_clip_textual_encoder

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

Source from the content-addressed store, hash-verified

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
157class TestCLIPModel(parameterized.TestCase):

Callers

nothing calls this directly

Calls 9

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

Tested by

no test coverage detected