Train your JAX models using model.fit(...) in Keras 3

Google for DevelopersAbout 5 min readMay 22, 2025Watch original
THE SUMMARYAI-generated

Keras, JAX, and Model Training: A Deep Dive

Key Concepts: Keras ecosystem, JAX backend, model.fit, progressive disclosure of complexity, callbacks, multi-device distribution, device mesh, layout map, Keras Hub, pre-trained models, presets, task models (backbone, preprocessor, task head), fine-tuning, LoRA, Q-LoRA, Keras Recommenders, two-phase approach (retrieval, ranking), multitask models, distributed embedding layer, TPUs, sparse-core processors.

1. Introduction to Keras and JAX Integration

  • Main Topic: Overview of the Keras ecosystem and its integration with JAX for training neural networks.
  • Key Points:
    • Keras provides a comprehensive toolbox for deep learning tasks with JAX.
    • It supports various levels of abstraction, from simple model.fit loops to advanced model parallel distribution.
    • Keras works with TensorFlow and PyTorch, allowing backend selection via the KERAS_BACKEND environment variable.
    • Data pipelines (e.g., tf.data) can be used independently of the chosen backend (e.g., JAX).

2. Training with model.fit: Progressive Complexity

  • Main Topic: Using the model.fit API for training JAX models, showcasing different levels of customization.
  • Key Points:
    • Simple Workflow: Compile the model with an optimizer, loss function, and metrics, then call model.fit with training data, batch size, epochs, and validation data.
    • Customization with Callbacks: Use callbacks like EarlyStopping to modify the training loop based on validation loss.
    • Multi-Device Distribution: Utilize keras.distribution APIs to enable multi-device training.
      • Device Mesh: Create a grid of available devices, defining its shape and assigning names to each axis (e.g., "data," "model" in a 2x4 grid).
      • Layout Map: Specify how to shard the model's weights across the device mesh.
      • Distribution Strategy: Configure the training process according to the sharding configuration.
    • Overriding Core Training Steps: Customize the training loop by overriding methods like compute_loss and compute_updates.

3. Keras Hub: Pre-trained Models and Presets

  • Main Topic: Introduction to Keras Hub, a model garden with pre-trained models and seamless integrations.
  • Key Points:
    • Keras Hub offers over 211 model variants based on 44+ state-of-the-art architectures (e.g., Gemma, Llama, Qwen, Stable Diffusion).
    • Models are ready to use with JAX, TensorFlow, or PyTorch.
    • Integrations with Kaggle and Hugging Face provide access to models, code snippets, and community guides.
    • Presets: Neatly packaged directories containing everything needed to load a pre-trained model (config, checkpoints, extra files).
      • Presets can be loaded using built-in identifiers (e.g., gemma_2b_en), Kaggle URLs, Hugging Face URLs, or local paths.
      • Custom models can be built with custom configurations.
    • Keras Hub models work with over 200,000 checkpoints via Kaggle and Hugging Face integrations.

4. Keras Hub Model Architecture: Task Models, Backbones, and Preprocessors

  • Main Topic: Deep dive into the underlying structure of Keras Hub models.
  • Key Points:
    • Task Models: High-level entry points that combine pre-processing and modeling into one class (e.g., causal LLM, image classifier, text classifier).
    • Backbone: Responsible for mapping preprocessed inputs to the model's latent space. Contains the architecture and pre-trained parameters.
    • Preprocessor: Converts raw data (images, text, audio) into a format the model can understand (tensors). Includes tools like tokenizers, audio converters, and image converters.
    • All components (backbone, preprocessor, tokenizer) can be instantiated using the from_preset method.

5. Fine-tuning Keras Hub Models: LoRA and Q-LoRA

  • Main Topic: Techniques for fine-tuning Keras Hub models efficiently.
  • Key Points:
    • Models can be fine-tuned using standard model.fit training loops.
    • LoRA (Low-Rank Adaptation): Reduces the number of trainable parameters by freezing the original model and training only two smaller matrices.
      • Enable LoRA by calling enable_lora and specifying the rank on the backbone.
    • Q-LoRA (Quantized LoRA): Adds quantization to LoRA, converting model data into smaller, more efficient formats.
      • Quantize the weights, then enable LoRA as usual.
    • Trained model weights can be uploaded to Kaggle and Hugging Face using keras_hub.upload_preset.

6. Keras Recommenders: Building Recommendation Systems

  • Main Topic: Introduction to Keras Recommenders, a library for building state-of-the-art recommendation systems.
  • Key Points:
    • Two-Phase Approach:
      • Retrieval Phase: Efficiently filters the item pool to identify a manageable set of candidates.
      • Ranking Phase: Uses a more intensive scoring method to select and order the most relevant items.
    • Multitask Models: Combine retrieval and ranking tasks, often outperforming single-task models due to transfer learning.
    • Keras Recommenders provides specialized layers, losses, and metrics for training recommender models.
    • Components are backend-agnostic, supporting JAX, TensorFlow, and Torch.
    • Examples and code for standard architectures are provided.

7. Example: Building a Sequential Recommender Model (GRU4Rec)

  • Main Topic: Practical example of building a sequential recommender model using Keras Recommenders.
  • Key Points:
    • Sequential recommenders leverage a user's past interactions to make predictions.
    • Two-Tower Architecture:
      • Query Tower: Represents the user, using a Gated Recurrent Unit (GRU) to process the sequence of past items.
      • Candidate Tower: Represents the items, using a simple embedding layer.
    • Retrieval Layer: Performs the actual recommendations.
    • Compute Loss Method: Overridden to calculate the affinity scores between the query tower output and the candidate tower output.
    • Categorical Cross-Entropy Loss: Used to quantify how far the affinity scores are from a perfect match.
    • The model is trained using model.fit with user-item interaction data.

8. Distributed Embedding Layer for Large-Scale Recommendations

  • Main Topic: Using the distributed embedding layer to handle large embedding tables in recommender systems.
  • Key Points:
    • Embedding tables can become too large for a single processing unit.
    • Google's TPUs with sparse-core processors are optimized for handling large embedding models.
    • The distributed embedding layer automatically shards the embedding tables across TPUs and sparse-core processors.
    • All embedding lookups are performed in a single pass for maximum performance.
    • Configuration:
      • Specify the dimensionality of each embedding table.
      • Define the combiner function (how to reduce embeddings obtained from looking up a sequence of items).
      • Configure features and connect them with the embedding tables.
    • The distributed embedding layer is available in Keras RS for both JAX and TensorFlow backends.

9. Conclusion

  • Main Takeaways: The Keras ecosystem provides a comprehensive set of tools for building and customizing machine learning models, particularly with JAX. From simple training loops to advanced distributed training and pre-trained models, Keras offers a progressive and flexible approach to deep learning. Keras Hub and Keras Recommenders extend the ecosystem with specialized tools for specific domains, simplifying the development of state-of-the-art models.

AI summaries can miss context or contain errors. Check important details against the original video.

Go a little deeper.

Have a question about this video? Load its transcript to open the video chat.