mirror of
https://github.com/google-deepmind/deepmind-research.git
synced 2026-10-06 06:59:41 +08:00
[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
This commit is contained in:
committed by
Saran Tunyasuvunakool
parent
3257aa3833
commit
f6da52cb38
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user