(self, x: jax.Array)
| 51 | ) |
| 52 | |
| 53 | def encode(self, x: jax.Array) -> jax.Array: |
| 54 | x = self.input_embedding_table[(x, )] |
| 55 | x *= jnp.sqrt(self.embed_dim).astype(x.dtype) |
| 56 | return x |
| 57 | |
| 58 | def decode(self, x: jax.Array) -> jax.Array: |
| 59 | return jnp.dot(x, self.input_embedding_table.T) |
no test coverage detected