Repository navigation
Patch release commits for 3.13.1 - #22005
Conversation
Merely importing keras currently triggers this warning with NumPy 2. ``` keras/src/export/tf2onnx_lib.py:8: FutureWarning: In the future `np.object` will be defined as the corresponding NumPy scalar. ``` Only patch NumPy if and when needed.
The signature of `check_is_flash_attention` changed with JAX 0.7.0. In addition to `query` and `key`, a positional argument of `value` was added. This was not caught as our kokoro tests use JAX 0.6.2 because it's the last version that supports Python 3.10. This change was tested here: keras-team#21957
This is a follow-up of keras-team#21674 This PR changed the signature of `make_tf_tensor_spec` from `(x)` to `(x, dynamic_batch=True)`, thereby adding the ability to make the batch size dynamic. This PR also adds `_get_save_spec(self, dynamic_batch=True)` which uses `make_tf_tensor_spec` and forwards the `dynamic_batch` argument. However, the default before this change for other export (SavedModel, ONNX) was to keep the batch size untouched. In particular, when a user manually provides an `input_signature` to [`ExportArchive.add_endpoint`](https://github.com/keras-team/keras/blob/master/keras/src/export/saved_model.py#L362), we should honor. The user controls whether the batch size is dynamic or not in the `input_signature`. This PR changes the default of `make_tf_tensor_spec` back to `dynamic_batch=False` to revert SavedModel and ONNX exports to the previous behavior. Also removed call to `return super()._get_save_spec(dynamic_batch)` which can never succeed as `TFLayer` is a top level class (ignoring the auto-tracking stuff).
Summary of ChangesHello @sachinprasadhs, I'm Gemini Code Assist1! I'm currently reviewing this pull request and will post my feedback shortly. In the meantime, here's a summary to help you and other reviewers quickly get up to speed! This pull request introduces Keras version 3.13.1, a patch release focused on addressing several critical bugs. It includes fixes for CUDNN flash attention compatibility with recent JAX versions, resolves regressions in the model export process, and enhances NumPy compatibility for ONNX export to prevent import-time warnings. These changes aim to improve the stability and robustness of the Keras framework. Highlights
🧠 New Feature in Public Preview: You can now enable Memory to help Gemini Code Assist learn from your team's feedback. This makes future code reviews more consistent and personalized to your project's style. Click here to enable Memory in your admin console. Using Gemini Code AssistThe full guide for Gemini Code Assist can be found on our documentation page, here are some quick tips. Invoking Gemini You can request assistance from Gemini at any point by creating a comment using either
Customization To customize Gemini Code Assist for GitHub experience, repository maintainers can create a configuration file and/or provide a custom code review style guide (such as PEP-8 for Python) by creating and adding files to a Limitations & Feedback Gemini Code Assist may make mistakes. Please leave feedback on any instances where its feedback is incorrect or counter productive. You can react with 👍 and 👎 on @gemini-code-assist comments. If you're interested in giving your feedback about your experience with Gemini Code Assist for Github and other Google products, sign up here. You can also get AI-powered code generation, chat, as well as code reviews directly in the IDE at no cost with the Gemini Code Assist IDE Extension. Footnotes
|
There was a problem hiding this comment.
Code Review
This pull request includes a set of patches for the 3.13.1 release, addressing a warning on import, a CUDNN flash attention issue with JAX, and an export regression. The changes appear correct and effectively resolve the described issues. The version bump is appropriate for a patch release. I've added a couple of suggestions to improve code quality and robustness in the JAX backend.
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## r3.13 #22005 +/- ##
========================================
Coverage ? 82.65%
========================================
Files ? 588
Lines ? 61266
Branches ? 9607
========================================
Hits ? 50642
Misses ? 8141
Partials ? 2483
Flags with carried forward coverage won't be shown. Click here to find out more. ☔ View full report in Codecov by Sentry. 🚀 New features to boost your workflow:
|
Remove NumPy warning with NumPy >= 2. #21949 (import keras always prints a warning)
Fix CUDNN flash attention for JAX > 0.6.2. #21970 (CUDNN flash attention broken with JAX > 0.6.2)
Do not always make batch size dynamic during export. #21944 (export regression)