
推理引擎大模型【免费下载链接】FlexGenRunning large language models on a single GPU for throughput-oriented scenarios.项目地址https://gitcode.com/gh_mirrors/fl/FlexGen点击查看免费下载本文是一份以 Hugging Face Flax/JAX Community Week 为核心场景的端到端技术指南。它完整覆盖了在本地与 TPUv3-8 上搭建 JAX/Flax Transformers 开发环境、理解 Flax 无状态设计哲学、用示例脚本完成前向推理与训练循环、通过 Hub 组织团队协作以及用 Widgets / Streamlit / Gradio 构建成果 Demo 的完整工作流。读完本文你将掌握一套可复制的「项目定义 → 环境安装 → 模型训练 → 协作发布 → Demo 展示」实战方法并能直接参考当前仓库中 jax-projects 目录 下的真实研究项目落地。活动全景Flax/JAX Community Week 的组织方式Flax/JAX Community Week 的目标是让计算密集型的 NLP 与 CV 项目如预训练 BERT、GPT2、CLIP、ViT对更广泛的工程师与研究者群体变得切实可行。活动期间参与团队可获得 Google Cloud 提供的TPUv3-8免费访问权限并用 JAX/Flax 完成一个 NLP 和/或 CV 项目。整个活动按以下节奏推进原始时间表见文档项目提案与组队期官方公告后参与者先提出项目创意再围绕最受关注的项目组成 35 人团队技术讲座期由 Google、Hugging Face 及开源社区科学家与工程师开设 JAX/Flax、TPU、Transformers、CV 与 NLP 主题讲座训练冲刺期每个团队获得 TPUv3-8 的访问权在约 10 天内完成项目训练成果评审期团队提交 Demo由评委会评审并评出前三名。所有关键沟通在内部 Slack 频道#flax-jax-community-week中进行包括公告、训练脚本发布、TPU 访问说明等。项目相关的技术问题则被引导至公开论坛分类以便问答沉淀为社区长期资产——这正是公开协作精神的体现问题不只在 Slack 里解决而是以可检索、可复用、避免重复回答的方式沉淀下来。项目策划与团队协作方法论如何提出项目并组队组织方提供了默认项目清单但强烈鼓励参与者提交自己的项目创意详细提案流程见 HOW_TO_PROPOSE_PROJECT.md。组队遵循一套开放的协作规则在论坛的项目分类下浏览项目创意用 ❤️ 表达兴趣直接在线程中留下评论、改进建议与细节疑问表达加入意愿时简要说明你的背景、时区、加入动机、可贡献方向与对项目的愿景项目热度足够后组织者会在线程中正式宣告团队成立并记录团队名与成员一个项目可组建多个团队参与者也可加入多个团队虽建议专注单一项目特殊情况下可调整团队。好项目的六个定义要素一个定义良好的项目应当让全体成员对以下问题达成共识任务模型将被训练来解决什么任务如摘要生成模型使用什么架构如 T5-small数据集训练数据来源如 CNN/Daily Mail训练脚本需要编写或适配哪类脚本产出项目的期望结果工作流数据 → 训练 → 评估 → Demo 的推进顺序并为各步骤设定截止时间。文档给出的示例是任务为摘要、模型为t5-small、数据为 CNN/Daily Mail、脚本为run_summarization_flax.py、产出为可摘要新闻的 T5 模型。一个项目不必所有代码从零编写但必须清楚数据集如何获取、训练脚本如何实现。工作量分配与期望管理团队效率的关键在于坦诚沟通预期参与度。建议经验丰富、积极性高的成员主导分工并在必要时接手他人任务。工作量通常可按下述维度切分数据预处理加载数据并格式化为模型可用的形态数据分词 / Data Collator将样本加工为 token 或图像模型配置定义模型的代码模型前向传播确保输入输出正确损失函数定义将各部件组装进训练脚本。由于上述步骤存在依赖关系文档建议先用预期格式的假数据启动模型前向传播调试不必等数据预处理完成。同时期望必须务实10 天 TPUv3-8 计算资源不足以预训练 110 亿参数的 T5也不适合从零手写整套建模、数据与训练代码。调试与协作小技巧尽量公开协作重要讨论发在论坛模型与训练日志托管在 Hub 上以享受版本控制缩短调试周期用datasets[train] datasets[train].select(range(1000))只加载前 1000 条样本跑通训练脚本而不是处理全量数据集也可利用 Datasets 的流式加载卡住时主动求助使用公开 Slack 频道或论坛分类。环境安装本地与 TPU VM 双端配置文档建议同时在本地机器与 TPU VM 上安装全部依赖本地用于快速原型与调试TPU 用于实际训练。核心依赖为 JAX、Flax、Optax、Transformers、Datasets 五个库并强烈建议全部安装在 Python 虚拟环境中。本地安装流程# 创建并激活虚拟环境 python3 -m venv your-venv-name source ~/your-venv-name/bin/activateFork Transformers 仓库克隆到本地并添加 upstream 远端创建开发分支git clone https://github.com/your Github handle/transformers.git cd transformers git remote add upstream https://github.com/huggingface/transformers.git git checkout -b a-descriptive-name-for-my-project在虚拟环境中以可编辑模式安装 Flax 依赖会自动带入flax、jax、optaxpip install -e .[flax]若环境中已有 transformers先pip uninstall transformers再以-e可编辑模式重装。从源码安装 Datasets以享受活动期间的最新更新git clone https://github.com/huggingface/datasets.git cd datasets pip install -e .[streaming]安装验证以下脚本同时验证 Transformers 的 Flax 模型、Datasets 的流式加载与 JAX 是否就绪前提是transformers与datasets均从 main 分支源码安装流式功能才可用from transformers import FlaxRobertaModel, RobertaTokenizerFast from datasets import load_dataset import jax dataset load_dataset(oscar, unshuffled_deduplicated_en, splittrain, streamingTrue) dummy_input next(iter(dataset))[text] tokenizer RobertaTokenizerFast.from_pretrained(roberta-base) input_ids tokenizer(dummy_input, return_tensorsnp).input_ids[:, :10] model FlaxRobertaModel.from_pretrained(julien-c/dummy-unknown) # 前向传播应返回 FlaxBaseModelOutputWithPooling 对象 model(input_ids)TPU VM 安装流程极重要同一时刻只能有一个进程访问 TPU 核心。若多人同时连接 TPU会抛出libtpu.so already in used by another process. Not attempting to load libtpu.so in this process.错误。因此文档建议每个成员可创建自己的虚拟环境但重训练进程应由一人运行且配置 TPUv3-8 时轮流验证 JAX 是否装好。在虚拟环境之外TPU 端还需两步pip install requests pip install jax[tpu]0.2.16 -f https://storage.googleapis.com/jax-releases/libtpu_releases.html该命令可能报invalid command bdist_wheel之类的构建错误但 JAX 通常仍会正确安装不必惊慌。验证 JAX 是否识别 TPUimport jax jax.device_count()在 TPUv3-8 VM 上应显示88 个 TPU 核心。TPU 端同样需要 Fork/克隆 Transformers 并以pip install -e .[flax]安装再源码安装 Datasets步骤与本地端一致可复用上文命令。JAX/Flax 快速入门生态与基础概念JAX 是 Autograd 与 XLA 的结合体面向高性能数值计算与机器学习研究提供对 PythonNumPy 程序的可组合变换微分、向量化、并行化、JIT 编译到 GPU/TPU 等。Flax 是构建于 JAX 之上的高性能神经网络库设计目标是与grad、pmap等 JAX 变换无缝配合给予用户对训练代码的完全控制。在 examples/flax/README.md 中给出了更凝练的概括JAX 通过jit将纯函数追踪并编译为 GPU/TPU 上高效融合的加速代码并支持grad任意梯度、pmap多设备并行、remat梯度检查点、vmap自动高效向量化、pjit自动切分的模型并行等变换且这些变换可任意组合——例如vmap(grad(f))即可高效计算逐样本梯度。Flax 基于这些能力提供以 Python dataclass 定义的模块抽象其提升式JAX 变换vmap、remat允许任意嵌套变换与模块。同一份 JAX/Flax 代码可无改动地运行在 CPU、GPU 与 TPU 上因为 JAX 由 XLA 编译器支撑所有 XLA 兼容设备都执行同一套编译产物。Transformers 中的 Flax 模型生态在仓库的 src/transformers/models 目录下可以逐一找到 Flax 实现的模型文件印证了文档中列出的支持范围文档撰写时点与仓库快照以仓库实际内容为准BERT、RoBERTa、ELECTRA —— 掩码语言建模 / 分类 / 长文本问答GPT2、OPT、GPT-Neo —— 因果语言建模T5、BART、MBART —— 摘要与 Seq2SeqCLIP、ViT —— 视觉与图文双编码BigBird —— 稀疏注意力的长序列建模Wav2Vec2 —— 语音表征预训练。对应的训练脚本同样在仓库中可查examples/flax因果语言建模GPT2run_clm_flax.py掩码语言建模BERT、RoBERTa、ELECTRA、BigBirdrun_mlm_flax.py掩码 Seq2Seq 预训练T5run_t5_mlm_flax.py文本分类GLUErun_flax_glue.py摘要 / Seq2SeqBART、MBART、T5run_summarization_flax.py问答SQuADrun_qa.py命名实体识别run_flax_ner.py图像分类run_image_classification.py。注意目前 JAX/Flax 后端没有 Trainer 抽象所有示例都显式编写训练循环——这也正是本文后续要深入讲解训练循环的原因。Flax 设计哲学有状态 vs 无状态理解 Flax 与 PyTorch 的根本差异是正确使用 Transformers Flax API 的前提。核心区别在于stateless无状态vs stateful有状态以及immutable不可变vs mutable可变。PyTorch 的有状态模型JAX 的大多数变换尤其是jax.jit要求被变换的函数无副作用副作用只在编译期执行一次之后所有调用都会复用编译期捕获的副作用而非真实状态。因此 PyTorch 式的设计权重存在模型实例内部model(input_ids)隐式使用内部权重天然与 JIT 冲突。文档中的示例模型包含注意力层的key_proj、value_proj、query_proj与输出层logits_proj。PyTorch 版本将权重存为torch.nn.Linear对象挂在类属性上class ModelPyTorch: def __init__(self, config): self.config config self.key_proj torch.nn.Linear(config) self.value_proj torch.nn.Linear(config) self.query_proj torch.nn.Linear(config) self.logits_proj torch.nn.Linear(config)实例化即分配权重内存可通过model_pytorch.key_proj.weight.data访问前向传播只需sequences model_pytorch(input_ids)。这是有状态的即使输入input_ids不变输出sequences也可能因内部状态权重改变而变化——输出不只依赖输入还依赖对象内部状态。Flax 的无状态模型等价的 Flax 版本将线性层声明为flax.linen.Dense对象。但实例化不会分配权重内存self.key_proj只是声明对输入执行线性变换并期待一个名为key_proj的权重作为输入。因此前向传播必须同时传入输入与权重字典# 先通过 dummy 输入显式初始化出初始权重 state state model_flax.init(rng, dummy_input_ids) # 再携带 state 执行前向传播 sequences model_flax.apply(state, input_ids)这是无状态的输出只由输入input_ids与state字典完全决定。随之而来的逻辑推论是 Flax 模型不可变若调用model_flax能改变自身则两次相同输入可能得到不同结果破坏无状态性。Transformers 如何包装无状态模型FlaxPreTrainedModel读者可能会问为什么日常使用FlaxBertModel时并不总是显式传参数因为 FlaxPreTrainedModel 抽象类把这一层隐藏了起来使 Flax 模型拥有与 PyTorch / TensorFlow 相似的用户 API。从源码可见类属性包括config_class配置类、base_model_prefix基础模型属性名、main_input_name主输入名NLP 多为input_ids视觉为pixel_values语音为input_values构造函数接收config、Flaxmodule、input_shape、seed、dtype默认jnp.float32与_do_init在_do_initTrue时通过self.init_weights(self.key, input_shape)自动完成随机初始化modeling_flax_utils.pyparams属性property暴露已初始化/加载的参数支持model.params读取与赋值并校验参数键是否齐全提供to_bf16/to_fp32/to_fp16方法modeling_flax_utils.py可在 TPU 上将参数显式转为bfloat16做全半精度训练或节省显存在 GPU 上转float16转换通过_cast_floating_to配合 mask 实现可选择性跳过某些参数如 LayerNorm 的 bias/scale。其工作方式如下init_weights接收期望输入形状与PRNGKey及初始化所需的其他参数调用module.init传入随机样本得到指定dtypefp32、bf16 等的初始权重。该过程在创建模型实例时即触发因此model FlaxBertModel(config)之后权重已就绪。__call__定义前向传播接收模型输入与可选参数不传参数时使用已初始化/加载的model.params内部调用module.apply完成实际计算即output model(inputs, paramsparams)。动手实现一个 Flax 预训练模型子类文档用一个两层 MLP 展示了完整模式。首先定义 Flax 模块声明层结构与计算import flax.linen as nn import jax.numpy as jnp class MLPModule(nn.Module): config: MLPConfig dtype: jnp.dtype jnp.float32 def setup(self): self.dense1 nn.Dense(self.config.hidden_dim, dtypeself.dtype) self.dense2 nn.Dense(self.config.hidden_dim, dtypeself.dtype) def __call__(self, inputs): hidden_states self.dense1(inputs) hidden_states nn.relu(hidden_states) hidden_states self.dense2(hidden_states) return hidden_states接着继承FlaxPreTrainedModel定义预训练模型包装类from transformers.modeling_flax_utils import FlaxPreTrainedModel class FlaxMLPPreTrainedModel(FlaxPreTrainedModel): config_class MLPConfig base_model_prefix model module_class: nn.Module None def __init__(self, config, input_shape(1, 8), seed0, dtypejnp.float32, **kwargs): module self.module_class(configconfig, dtypedtype, **kwargs) super().__init__(config, module, input_shapeinput_shape, seedseed, dtypedtype) def init_weights(self, rng, input_shape): inputs jnp.zeros(input_shape, dtypei4) params_rng, dropout_rng jax.random.split(rng) rngs {params: params_rng, dropout: dropout_rng} params self.module.init(rngs, inputs)[params] return params def __call__(self, inputs, paramsNone): params {params: params or self.params} outputs self.module.apply(params, jnp.array(inputs)) return outputs最后一行定义具体模型类class FlaxMLPModel(FlaxMLPPreTrainedModel): module_class FlaxMLPModule要点FlaxMLPModel实例本身有状态持有全部参数而底层 Flax 模块FlaxMLPModule依旧无状态训练时随时可以显式传参给模型以完全契合 JAX 变换。另一个显著差异是labels的处理PyTorch 允许把labels直接传给前向计算损失然后.backward()反向传播而 Flax 的前向函数不允许传入labels——因为反向传播必须通过jax.grad/jax.value_and_grad变换损失函数来获得所有参数的梯度这种变换无法在前向内部自动完成。因此所有训练相关代码都与建模代码解耦显式写在训练脚本中。这正是 examples/flax 下所有脚本没有 Trainer 抽象的原因。实战一前向推理使用FlaxRobertaModel演示加载、保存与推理。为获得最佳性能用jax.jit编译函数注意 JAX 编译器在输入形状变化时需要重新编译因此统一使用paddingmax_length将样本填充到固定静态形状如 128避免频繁重编译from transformers import FlaxRobertaModel, RobertaTokenizerFast import jax tokenizer RobertaTokenizerFast.from_pretrained(roberta-base) inputs tokenizer(JAX/Flax is amazing , paddingmax_length, max_length128, return_tensorsnp) model FlaxRobertaModel.from_pretrained(julien-c/dummy-unknown) jax.jit def run_model(input_ids, attention_mask): # 返回 FlaxBaseModelOutputWithPooling 对象 return model(input_ids, attention_mask) outputs run_model(**inputs)实战二完整训练循环以FlaxGPT2ForCausalLM为例Flax 训练循环由四个部件组成损失函数接收参数与输入前向传播后返回损失变换用jax.grad或jax.value_and_grad变换损失函数以获取梯度优化器用梯度更新参数训练步组合损失函数与优化器更新完成前向与反向返回更新后的参数。import jax import jax.numpy as jnp from transformers import FlaxGPT2ForCausalLM from flax.training.common_utils import onehot model FlaxGPT2ForCausalLM(config) def cross_entropy(logits, labels): return -jnp.sum(labels * jax.nn.log_softmax(logits, axis-1), axis-1) # 定义前向 损失计算注意显式传 params def compute_loss(params, input_ids, labels): logits model(input_ids, paramsparams, trainTrue) num_classes logits.shape[-1] loss cross_entropy(logits, onehot(labels, num_classes)).mean() return loss # 变换损失函数以获得梯度 grad_fn jax.value_and_grad(compute_loss) # 用 optax 初始化优化器 import optax params model.params tx optax.sgd(learning_rate3e-3) opt_state tx.init(params) # 单步训练前向 反向 def _train_step(params, opt_state, input_ids, labels): loss, grads grad_fn(params, input_ids, labels) updates, opt_state tx.update(grads, opt_state) updated_params optax.apply_updates(params, updates) return updated_params, opt_state, loss train_step jax.jit(_train_step) # 训练循环 for i in range(10): params, opt_state, loss train_step(params, opt_state, input_ids, labels)注意每一步都把params与opt_state传入并回收更新后的版本——这正是无状态模型的外部状态管理模式。训练完成后保存model.save_pretrained(awesome-flax-model, paramsparams)由于 JAX 由 XLA 支撑同一份代码可以不加改动地在 CPU、GPU、TPU 上运行。用 Hub 组织团队训练一个完整实战案例Hub 是团队协作的核心载体每个团队可在组织下创建带 git 版本控制的模型仓库享受便捷协作成员均有写权限、集成 git 版本控制代码与大型模型文件统一追踪、轻松分享与内置 TensorBoard上传的 trace 自动展示。创建仓库并推送配置文档以在低资源语言 Alemannicals上预训练 RoBERTa 为例登录 Hub 后在组织如flax-community下创建公开仓库roberta-base-als然后本地操作huggingface-cli login git clone https://huggingface.co/flax-community/roberta-base-als cd ./roberta-base-als在 Python shell 中生成并保存配置from transformers import RobertaConfig config RobertaConfig.from_pretrained(roberta-base) config.save_pretrained(./)上传配置git add . git commit -m add config git push训练低资源语言分词器用 OSCAR 的unshuffled_deduplicated_als数据集与tokenizers库训练 ByteLevel BPE 分词器from datasets import load_dataset from tokenizers import ByteLevelBPETokenizer dataset load_dataset(oscar, unshuffled_deduplicated_als, splittrain) tokenizer ByteLevelBPETokenizer() def batch_iterator(batch_size1000): for i in range(0, len(dataset), batch_size): yield dataset[i: i batch_size][text] tokenizer.train_from_iterator(batch_iterator(), vocab_size50265, min_frequency2, special_tokens[ s, pad, /s, unk, mask, ]) tokenizer.save(./tokenizer.json)启动训练将官方 MLM 脚本复制进仓库确保训练所用代码都被版本控制追踪然后运行./run_mlm_flax.py \ --output_dir./ \ --model_typeroberta \ --config_name./ \ --tokenizer_name./ \ --dataset_nameoscar \ --dataset_config_nameunshuffled_deduplicated_als \ --max_seq_length128 \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --learning_rate3e-4 \ --warmup_steps1000 \ --overwrite_output_dir \ --num_train_epochs8 \ --push_to_hub该数据集很小整个命令预计在 5 分钟内完成。--push_to_hub标志让模型权重与 TensorBoard trace 自动上传到 Hub训练指标直接展示在模型页面的 Training metrics 标签页。由于仓库自带 git 版本控制与 git-lfs模型权重等大文件也可轻松上传与变更。上传任意 Flax 模型非 Transformers 模型也适用若不使用 Transformers 模型huggingface_hub同样支持用几行代码上传任何 JAX/Flax 模型需huggingface_hub 0.0.13from flax import serialization from jax import random from flax import linen as nn from huggingface_hub import Repository model nn.Dense(features5) key1, key2 random.split(random.PRNGKey(0)) x random.normal(key1, (10,)) params model.init(key2, x) bytes_output serialization.to_bytes(params) repo Repository(flax-model, clone_fromflax-community/flax-model-dummy, use_auth_tokenTrue) with repo.commit(My cool Flax model :)): with open(flax_model.msgpack, wb) as f: f.write(bytes_output)TPU VM 连接gcloud 实战获得 TPU 访问权后你会收到两封邮件其一授予hf-flax项目中 Community Week Participants 角色其二或多封若参与多个项目给出团队 TPU 的名称与 zone。即使 Cloud Console 上无法可视化查看hf-flax项目并提示 You dont have sufficient permission to view this page也属预期现象。连接步骤安装 Google Cloud SDK设置账号须与报名邮箱一致gcloud config set account your-email-address设置项目若邮箱关联多个 gcloud 项目gcloud config set project hf-flax认证gcloud auth loginSSH 进入 TPU VMzone取邮件中的europe-west4-a或us-central1-atpu-name取邮件中的 TPU 名gcloud alpha compute tpus tpu-vm ssh tpu-name --zone zone --project hf-flax进入后按前文TPU VM 安装流程安装依赖并注意 JAX 的 TPU 专用安装与单进程访问限制。也可通过 VS Code 等 IDE 连接 TPU VM。构建成果 Demo 的三条路径所有团队必须提交 Demo。文档提供了三条路线路径一Hugging Face WidgetsHub 内置 15 种开箱即用的推理组件按领域可分为NLPConversational对话、Feature Extraction特征提取、Fill Mask掩码预测、Question Answering抽取式问答、Sentence Similarity句相似度、Summarization摘要、Table Question Answering表格问答、Text Generation文本生成、Token ClassificationNER/POS、Zero-Shot Classification零样本分类等SpeechAudio to Audio音频分离/增强、Automatic Speech Recognition语音转文字、Text to Speech文字转语音ImageImage Classification图像分类等另有多项 WIP 组件零样本图像分类、图像描述、文生图、视觉问答。Widgets 全部开源也欢迎通过提 Issue 的方式提议与实现新组件。当使用场景恰好落在某类组件内时这是成本最低的展示方式。路径二Streamlit 演示当使用非主流库或特殊应用时Streamlit 是更灵活的纯 Python 方案。huggingface_hub帮助从模型仓库加载文件pip install huggingface_hubfrom huggingface_hub import hf_hub_download filepath hf_hub_download(flax-community/roberta-base-als, flax_model.msgpack)下载整个仓库可指定 revisionfrom huggingface_hub import snapshot_download local_path snapshot_download(flax-community/roberta-base-als)若使用 Transformers可直接加载模型与分词器from transformers import AutoTokenizer, AutoModelForMaskedLM tokenizer AutoTokenizer.from_pretrained(REPO_ID) model AutoModelForMaskedLM.from_pretrained(REPO_ID)路径三Gradio 演示Gradio 同样可快速为 Hugging Face 模型创建 GUI 并分享 Demo适合交互式展示。仓库内的社区周项目成果五个可参考的真实案例当前仓库 jax-projects 目录下沉淀了多个社区周期间的真实研究项目是上述工作流的直接产物与最佳范本1. BigBird 长文档问答big_bird在 Natural Questions 数据集上微调 BigBird 做长文档问答。BigBird 是基于稀疏注意力的 Transformer可将 BERT 类模型扩展到更长的序列。完整流程安装依赖 →python3 prepare_natural_questions.py下载并处理数据集约 100GB 磁盘、约 3 小时→python3 train.py启动训练TPUv3-8 上每轮约 4.5 小时两轮收敛→python3 evaluate.py评估序列长度至 4096评估脚本获得约 55.2 的 EM 分数。超参调优可用wandb sweep --projectbigbird sweep_flax.yaml配合wandb agent执行。配置见 bigbird_flax.py。2. 流式掩码语言建模dataset-streaming展示如何在 JAX/Flax 后端结合 Datasets 流式加载在单台 TPUv3-8 上预训练roberta-base英语10000 更新步全程无需下载完整数据集。相比默认run_mlm_flaxrun_mlm_flax_stream.py 新增 4 个训练设置num_train_steps更新步数num_eval_samples用于评估的训练样本数logging_steps训练损失记录频率eval_steps评估运行频率。典型运行参数含--adam_beta10.9、--adam_beta20.98、--per_device_train_batch_size128、--num_eval_samples5000、--logging_steps250、--eval_steps1000。仓库创建、git lfs track *tfevents*追踪 TensorBoard trace 等前置步骤与其 README 一一对应。3. 混合 CLIP 图文双编码器hybrid_clip用预训练文本与视觉编码器联合训练 CLIP 风格的图文双编码器将图像与描述映射到同一嵌入空间用于自然语言图像检索与零样本图像分类。FlaxHybridCLIP类可从任意文本/视觉编码器组合构造from modeling_hybrid_clip import FlaxHybridCLIP # 用预训练模型组合构造双编码器 model FlaxHybridCLIP.from_text_vision_pretrained(bert-base-uncased, openai/clip-vit-base-patch32) # 保存与加载 model.save_pretrained(bert-clip) model FlaxHybridCLIP.from_pretrained(bert-clip)若检查点是 PyTorch 格式可传text_from_ptTrue与vision_from_ptTrue自动转换加载。训练数据来自 MS-COCO8.2 万 图像、每图至少 5 条描述经数据整理脚本生成train_dataset.json/valid_dataset.json后用run_hybrid_clip.py训练其 README 注明了图像解码可预先完成以进一步提升数据加载性能。4. GPTNeo 模型并行训练model_parallel展示用 JAX 的pjit定义 GPTNeo 模型的PartitionSpecPyTree 描述切分方式实际切分由pjit自动处理。前置步骤颇具实战价值GPTNeo 词表大小为 50257需先将其调整到设备数的倍数如 50264再加载预训练权重并保存随后用run_clm_mp.py训练示例参数--dtype bfloat16、--block_size 1024、--learning_rate 4e-6、wikitext-2-raw-v1 数据集。5. Wav2Vec2 对比损失预训练wav2vec2run_wav2vec2_pretrain_flax.py演示 JAX/Flax 后端的 Wav2Vec2 预训练。其 README 特别强调大量训练参数可写入模型配置包括掩码分布mask_time_length、mask_time_prob、dropoutattention_dropout等、对比损失与多样性损失的权衡diversity_loss_weight、num_negatives等并推荐结合原始论文调整。示例用facebook/wav2vec2-base配置为mask_time_length10、mask_time_prob0.05、diversity_loss_weight0.1、num_negatives100、do_stable_layer_normTrue、feat_extract_normlayer特征提取器沿用return_attention_maskTrue训练数据用 LibriSpeechclean配置的train.100切分。项目评估评审标准与流程项目的评审依据四个维度Demo所有项目必须提交可演示成果形式开放可参考上文三条 Demo 构建路径技术难度涵盖复杂架构、优于现有模型的评估指标、低资源语言模型实现等方面社会影响期望项目产生积极社会影响如服务少数/弱势群体低资源语言、偏差公平与伦理议题或应对健康、气候等社会挑战创新性提出新颖应用或新思路的项目获得更高评价。评审流程为TPU 访问关闭 → 项目完成含 Demo→ 组织者初筛出前 15 名 → 评委会从中评选前三名并公布。评审团包括来自 Google Research 与 Hugging Face 的代表。这一流程提醒我们Demo 的可演示性、技术的含金量、社会价值与创新性缺一不可。总结从活动指南到可复用的工程方法论这份 Flax/JAX Community Week 指南本质上是一套完整的深度学习研究项目工作流通过公开、结构化的协作方式定义项目在本地与 TPU 双端搭建 JAX/Flax Transformers 环境借助FlaxPreTrainedModel抽象与显式训练循环在 TPU 上高效训练再以 Hub 的 git 版本控制与 TensorBoard 集成沉淀模型、日志与代码最终用 Widgets / Streamlit / Gradio 将成果转化为可演示、可评估的 Demo。仓库中的五个研究项目BigBird 长文档问答、流式 MLM 预训练、混合 CLIP、GPTNeo 模型并行、Wav2Vec2 预训练均为该方法论的真实落地。无论你是在低资源语言上预训练语言模型还是在长文档上微调问答系统这套从「项目定义」到「Demo 交付」的完整链路都可直接复用。赞分享推理引擎大模型【免费下载链接】FlexGenRunning large language models on a single GPU for throughput-oriented scenarios.项目地址https://gitcode.com/gh_mirrors/fl/FlexGen点击查看免费下载相关推荐Prometheus 配置文件 prometheus.yml 全解析scrape_configs 与全局设置一次讲透的完整清单Prometheus 配置文件 prometheus.yml 全解析scrape_configs 与全局设置一次讲透的完整清单 Prometheus普罗米修可观测性指标监控时序数据库告警开源项目 Mantle 亮点深度解析Objective-C模型层的革命性简化开源项目 Mantle 亮点深度解析Objective C模型层的革命性简化 引言Objective C模型开发的痛点 在iOS/macOS开发中处理JS推理引擎大模型Flax WMT 机器翻译示例全解析用 JAX/Flax 训练 Transformer 英德翻译模型并部署到 Cloud TPUFlax WMT 机器翻译示例全解析用 JAX/Flax 训练 Transformer 英德翻译模型并部署到 Cloud TPU 本指南以 Flax 官方仓库人工智能深度学习机器学习上一篇sd-scripts与IPEX集成Intel硬件上的极致性能优化下一篇jeecg-boot-starter 技术文档创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考