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 任务,预处理器负责:
- 添加起始和结束文本标记。
- 对标记序列进行填充并生成掩码。
- 生成用于训练的“期望输出”(即输入字符串向后移一位)。
训练与 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)尤为有用。
用户可以通过定义 DeviceMesh 和 LayoutMap 将模型拆分到多个 GPU 或 TPU 上。虽然大多数模型通过 get_layout_map(device_mesh) 提供了合理的默认值,用户仍可自定义布局映射以优化性能。例如,在 TPU v5e 上使用自定义布局映射可将 epoch 时间从 62 秒降低至 54 秒。
Sources
- Original“Llama 3.2 in Keras”