Keras 中的 Llama 3.2

Llama 3.2 已可在 Keras 中立即使用,支持加载标准的 Hugging Face 检查点(包括 safetensors),如有必要可进行即时转换。此集成使开发者能够在 Keras 生态系统中利用 Llama 3.2,受益于多后端灵活性和集成的训练工具。

支持 JAX、PyTorch 和 TensorFlow 的多后端

Keras 作为一个多后端建模库运行,使同一模型能够在 JAX、PyTorch 或 TensorFlow 上运行。后端通过在导入 Keras 之前设置环境变量来指定:

import os
os.environ["KERAS_BACKEND"] = "jax" # Options: "jax", "torch", or "tensorflow"

这种灵活性使用户能够使用带有 XLA 编译的 JAX,以获得优化的性能。

Keras-Hub 与模型集成

keras-hub(前称 KerasNLP 和 KerasCV)是 Keras 的预训练模型集合。它提供了流行模型的标准 Keras 实现,包括 Llama 3、Gemma、StableDiffusion 和 Segment Anything。

Llama 3.2 可以使用 keras_hub 中的 Llama3CausalLM 类加载:

from keras_hub import models.Llama3CausalLM
model = Llama3CausalLM.from_preset("hf://meta-llama/Llama-3.2-1B-Instruct", dtype="bfloat16")

“开箱即用” LLM 能力

Keras LLM 旨在通过将分词器直接集成到模型对象中来简化使用。这使得可以对原始字符串进行高级操作:

  • 生成: model.generate("Hi there!") 直接从字符串输入生成文本输出。
  • 训练: model.fit(strings) 直接在字符串列表或数据集上进行训练。

对话与指令微调

指令微调的变体,例如 Llama-3.2-1B-Instruct,支持使用特定标签格式的轮流对话。Llama 3.2 所需的格式包括 <|start_header_id|>system<|end_header_id|><|start_header_id|>user<|end_header_id|><|eot_id|> 等标签。格式化后,这些字符串可以直接传递给 model.generate()

低层模型访问

对于需要更多控制的用户,Keras 提供对底层组件的访问:

  • 分词器: 可通过 model.preprocessor.tokenizer 访问。它将文本转换为整数向量。
  • 骨干网络: 可通过 model.backbone 访问核心模型架构。

预处理器概念

Preprocessor(预处理器)在 Keras 中是用于数据转换的综合工具。对于 CausalLM 任务,预处理器负责:

  1. 添加起始和结束文本标记。
  2. 对标记序列进行填充并生成掩码。
  3. 生成用于训练的“期望输出”(即输入字符串向后移一位)。

训练与 Hub 集成

Keras 包含一个内置训练器,可通过 model.fit(ds) 访问。该训练器兼容 Keras 的特性,包括分布式训练、混合精度、量化以及 LoRA、QLoRA 等参数高效微调方法。

微调后的模型可以在本地使用 model.save_to_preset() 保存后,直接通过 keras_hub.upload_preset() 上传到 Hugging Face Hub。

分布式模型并行

Keras 通过 JAX 和 XLA 编译器提供了一条简化的高级模型并行路径,这对于单个加速器无法容纳的模型(例如 Llama 3.1 8B)尤为有用。

用户可以通过定义 DeviceMeshLayoutMap 将模型拆分到多个 GPU 或 TPU 上。虽然大多数模型通过 get_layout_map(device_mesh) 提供了合理的默认值,用户仍可自定义布局映射以优化性能。例如,在 TPU v5e 上使用自定义布局映射可将 epoch 时间从 62 秒降低至 54 秒。

Sources