Я столкнулся с некоторыми загадочными проблемами градиента NaN в процессе обучения модели, реализованной с помощью льна. С помощью jax_debug_nans я могу определить, что это происходит из-за градиента реализованного блока преобразователя, код которого представлен ниже, однако довольно сложно понять, где что-то идет не так. Трассировка стека ошибок, полученная с помощью jax_debug_nans:
Traceback (самый последний вызов — последний): Файл "/home/_/_/codes/_/_/main.py", строка 61, в app.run(основной) Файл "/home/_/_/_/envs/_/lib/python3.10/site-packages/absl/app.py", строка 308, в запуске _run_main(основной, аргументы) Файл "/home/_/_/_/envs/_/lib/python3.10/site-packages/absl/app.py", строка 254, в _run_main sys.exit(main(argv)) Файл "/home/_/_/codes/_/_/main.py", строка 53, в основном train_eval.train(FLAGS.config, FLAGS.workdir) Файл "/home/_/_/codes/_/_/_/train_eval.py", строка 143, в поезде (rng, обучающее_состояние), потеря = train_step_fn((rng, обучающее_состояние), обработанные_данные) Файл «/home/_/_/codes/_/_/_/training/losses.py», строка 304, в шаге_fn (потеря, (y_pred_mu, y_pred_sigma, mu_context, sigma_context, mu_tgt, sigma_tgt, mc_mean, kl_mean)), grad = grad_fn(step_rng, params, пакет) Файл «/home/_/_/codes/_/_/_/training/losses.py», строка 29, в анонимном_лоссе jax.vmap(partial_model_apply, in_axes=0)(data_x, data_y, data_x, data_y, context_mask, target_mask) Файл "/home/_/_/codes/_/_/_/training/losses.py", строка 24, в model.apply(переменные, x_context=x_ctx, y_context=y_ctx, x_target=x_tgt, \ Файл "/home/_/_/codes/_/_/_/models/model.py", строка 446, в __call__ v_star = self.query_specify_encode(x_context, y_context, context_mask, x_target) Файл "/home/_/_/codes/_/_/_/models/model.py", строка 408, в query_specific_encode v_star = self.qkv_to_v_star(q, k, v, ctx_mask) Файл "/home/_/_/codes/_/_/_/models/utils/nn.py", строка 152, в __call__ h = _scaled_dot_product_attention(qs, ks, vs, Файл «/home/_/_/codes/_/_/_/models/utils/nn.py», строка 34, в _scaled_dot_product_attention ws = softmax( Файл "/home/_/_/_/envs/_/lib/python3.10/site-packages/jax/_src/nn/functions.py", строка 352, в softmax return _softmax_deprecated (x, ось, где, начальный) Файл "/home/_/_/_/envs/_/lib/python3.10/site-packages/jax/_src/nn/functions.py", строка 377, в _softmax_deprecated результат = ненормализованный/jnp.sum(ненормализованный, ось, где=где, Keepdims=True) Файл "/home/_/_/_/envs/_/lib/python3.10/site-packages/jax/_src/numpy/array_methods.py", строка 791, в операции return getattr(self.aval, f"_{name}")(self, *args) Файл "/home/_/_/_/envs/_/lib/python3.10/site-packages/jax/_src/numpy/array_methods.py", строка 258, в deferring_binary_op вернуть двоичный_оп (* аргументы) jax._src.source_info_util.JaxStackTraceBeforeTransformation: FloatingPointError: недопустимое значение (nan), обнаруженное в jit (mul) Предыдущая трассировка стека является источником операции JAX, которая после преобразования JAX вызвала следующее исключение. Я провел обширное расследование и пришел к следующим выводам:
[*]Это не связано со слишком большой скоростью обучения (1e-4), кривая потерь и предел весов и градаций выглядят нормально (см. рис. в конце), также налагается глобальное ограничение нормы градации (так же, как быстрая Optax запустить блокнот) [*]Проблема возникает только из-за определенных данных в пакете, все остальные данные работают, я дважды проверил все данные и нет NaN/Inf вообще. [*]Я пытался либо использовать jax_softmax_custom_jvp, чтобы использовать не устаревший softmax, либо использовать config.update("jax_enable_x64", True), чтобы потенциально устранить любую проблему, связанную с числовыми значениями. точность, однако тут не повезло. [*]NaN Град встречается только в grad (не значение потери и параметры для расчета значения потери), более конкретно, градация NaN в основном встречается в блоках внимания и представлена в этом gist.
Блок внимания реализован как:
def _scaled_dot_product_attention(Q: jax.Array, K: jax.Array, V: jax.Array, Q_mask: jax.Array, K_mask: jax.Array) -> jax.Array: d_k = Q.shape[-1] ws = softmax( np.matmul(Q * Q_mask[..., Нет], rerange(K * K_mask[..., None], '... seq_length key_dim -> ... key_dim seq_length')) /d_k**0.5,\ где=K_mask, начальное=0.0, ) return np.matmul(ws, (V * K_mask[..., None])) класс MultiHeadCrossAttentionBlock(MultiHeadSelfAttentionBlock): """ Блок внимания корса """ def __call__(self, запросы: jax.Array, ключи: jax.Array, значения: jax.Array, keys_mask: Необязательно[jax.Array]) -> Массив: rerange_arg = ( "... (num_heads key_dim) -> num_heads... key_dim" ) qs = переставить( self.projs_q(запросы), переставить_арг, num_heads=self.heads_num, key_dim=self.key_dim, ) кс = переставить( self.projs_k(ключи), переставить_арг, num_heads=self.heads_num, key_dim=self.key_dim, ) vs = переставить( self.projs_v(значения), переставить_арг, num_heads=self.heads_num, key_dim=self.key_dim, ) h = _scaled_dot_product_attention(qs, ks, vs, Q_mask = np.ones(shape=(queries.shape[:-1])).astype(np.bool_), K_mask=keys_mask) # [num_heads, ..., H] # объединение и проекция h = self.proj( np.squeeze( np.concatenate( np.split(h, index_or_sections=self.heads_num, axis=0), axis=-1 ), ось=0, ) ) # proj: [num_heads, ..., H] -> [..., num_heads * H] -> [..., H] вернуть ч Я почти выбился из головы и застрял на некоторое время, буду очень признателен за любые подсказки по этому вопросу!
Дополнительная информация:

