From f6da52cb38f549ace94ec5757b2f78023bf9ba2a Mon Sep 17 00:00:00 2001 From: Peter Hawkins Date: Wed, 6 Oct 2021 14:57:06 +0100 Subject: [PATCH] [JAX] Replace uses of deprecated `jax.ops.index_update(x, idx, y)` APIs with their up-to-date, more succinct equivalent `x.at[idx].set(y)`. The JAX operators: jax.ops.index_update(x, jax.ops.index[idx], y) jax.ops.index_add(x, jax.ops.index[idx], y) ... have long been deprecated in lieu of their more succinct counterparts: x.at[idx].set(y) x.at[idx].add(y) ... This change updates users of the deprecated APIs to use the current APIs, in preparation for removing the deprecated forms from JAX. The main subtlety is that if `x` is not a JAX array, we must cast it to one using `jnp.asarray(x)` before using the new form, since `.at[...]` is only defined on JAX arrays. PiperOrigin-RevId: 401233414 --- wikigraphs/wikigraphs/model/transformer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/wikigraphs/wikigraphs/model/transformer.py b/wikigraphs/wikigraphs/model/transformer.py index 7362ff7..0f2f32a 100644 --- a/wikigraphs/wikigraphs/model/transformer.py +++ b/wikigraphs/wikigraphs/model/transformer.py @@ -331,7 +331,7 @@ def unpack_and_pad( idx = jnp.arange(total_size) idx += repeat_rows(jnp.arange(n_splits), split_sizes, total_size) * pad_size idx -= repeat_rows(cumsum, split_sizes, total_size) - out = jax.ops.index_update(out, idx, packed) + out = out.at[idx].set(packed) out = out.reshape([n_splits, pad_size] + out_shape[1:]) return out, masks