LWM:突破 100 万 Token 的开源大模型
开源模型 LWM 支持 100 万 token 上下文窗口,刷新开源模型纪录,大幅扩展长文本处理能力。
开源模型 LWM 支持 100 万 token 上下文窗口,刷新开源模型纪录,大幅扩展长文本处理能力。
大型世界模型(Large World Model, LWM)是一个通用的大上下文多模态自回归模型。它使用 RingAttention 在大规模多样化长视频和书籍数据集上进行训练,可以执行语言、图像和视频的理解与生成。
当前的语言模型在理解不易用文字描述的世界方面存在不足,并在处理复杂的长文本任务时表现有限。视频序列提供了语言和静态图像中不存在的宝贵时间信息,使其成为与语言联合建模的理想选择。这样的模型可以发展出对人类文本知识和物理世界的理解,从而实现更广泛的 AI 能力来辅助人类。然而,从数百万个 token 的视频和语言序列中学习会面临内存限制、计算复杂性和数据集有限等挑战。为了应对这些挑战,我们策划了大规模多样化视频和书籍数据集,利用 RingAttention 技术在长序列上可扩展地进行训练,并逐步将上下文大小从 4K 增加到 1M token。本论文做出以下贡献:(a) 最大上下文大小的神经网络:我们在长视频和语言序列上训练了最大上下文大小的 transformer 之一,在难度较大的检索任务和长视频理解方面设立了新的基准。(b) 克服视觉-语言训练挑战的解决方案,包括使用掩码序列打包来混合不同的序列长度、损失加权来平衡语言和视觉,以及基于模型生成的问答数据集用于长序列对话。(c) 一个高度优化的实现,具有 RingAttention、掩码序列打包和其他关键特性,用于在数百万长度的多模态序列上进行训练。(d) 完全开源的 7B 参数模型家族,能够处理超过 1M token 的长文本文档(LWM-Text、LWM-Text-Chat)和视频(LWM、LWM-Chat)。这项工作为在大规模长视频和语言数据集上进行训练奠定了基础,以发展对人类知识和多模态世界的理解,以及更广泛的能力。
LWM 可以在 1M 上下文中以高精度检索事实。
LWM 可以回答关于 1 小时 YouTube 视频的问题。
LWM 可以与图像进行对话。
LWM 可以从文本生成视频和图像。
本代码库在 Ubuntu 上得到支持,尚未在 Windows 或 macOS 上进行测试。我们建议使用 TPU 进行训练和推理,尽管也可以使用 GPU。在 TPU 上,代码使用 Jax 的 Pallas 进行了高度优化,可以在 RingAttention 的非常大上下文大小下实现高 MFU。在 GPU 上,代码基于 XLA,优化程度不如 TPU 版本。
使用以下方式安装所需依赖:
conda create -n lwm python=3.10
conda activate lwm
pip install -r gpu_requirements.txt
或使用以下方式设置 TPU VM:
sh tpu_requirements.sh
有仅语言版本和视频-语言版本,提供 32K、128K、256K 和 1M token 的上下文大小。视觉-语言模型仅在 Jax 中可用,仅语言模型在 PyTorch 和 Jax 中均可用。以下是可用模型的名称及其相应的上下文大小和能力:
使用 scan_query_chunk_size 和 scan_key_chunk_size 来控制自注意力分块计算中的块大小。使用 scan_mlp_chunk_size 来控制前馈网络分块计算中的块大小。使用 scan_attention=True 和 scan_mlp=True 来启用/禁用自注意力和前馈网络中的分块计算。
你可以使用 mesh_dim=dp, fsdp, tp, sp 来控制并行度和 RingAttention。这是一个由逗号分隔的 4 个整数的字符串,分别表示数据并行度、完全分片数据并行度、张量并行度和序列并行度。例如,mesh_dim='1,64,4,1' 表示 1 个数据并行、64 个完全分片数据并行、4 个张量并行和 1 个序列并行。mesh_dim='1,1,4,64' 表示 1 个数据并行、1 个完全分片数据并行、4 个张量并行和 64 个序列并行用于 RingAttention。
本部分提供了如何运行每个提供的脚本的说明。对于每个脚本,你可能需要填入脚本开头所述变量中的自己的路径和值。
要运行以下脚本,请使用 bash <script_name>.sh:
语言模型训练:bash scripts/run_train_text.sh
视觉-语言模型训练:bash scripts/run_train_vision_text.sh
单针评估(语言模型):bash scripts/run_eval_needle.sh
多针评估(语言模型):bash scripts/run_eval_needle_multi.sh
采样图像(视觉-语言模型):bash scripts/run_sample_image.sh
采样视频(视觉-语言模型):bash scripts/run_sample_video.sh
图像/视频理解(视觉-语言模型):bash scripts/run_vision_chat.sh
默认情况下,mesh_dim 参数将所有设备放在 tp(张量并行)上。对于较长的序列,你可能想要包含 sp,它是 mesh_dim 中的最后一个维度。
运行针评估时,你可能需要根据模型调整脚本中的 theta 和 max_sequence_length 参数。下面显示了每个模型的正确值。
填充脚本(run_sample_video.sh)的示例如下:
#! /bin/bash
export SCRIPT_DIR="$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd )"
export PROJECT_DIR="$( cd -- "$( dirname -- "$SCRIPT_DIR" )" &> /dev/null && pwd )"
cd $PROJECT_DIR
export PYTHONPATH="$PYTHONPATH:$PROJECT_DIR"
export llama_tokenizer_path="LargeWorldModel/LWM-Text-1M"
export vqgan_checkpoint="/path/to/ckpt/folder/vqgan"
export lwm_checkpoint="params::/path/to/ckpt/folder/params"
python3 -u -m lwm.vision_generation \
--prompt='Fireworks over the city' \
--output_file='fireworks.mp4' \
--temperature_image=1.0 \
--temperature_video=1.0 \
--top_k_image=8192 \
--top_k_video=1000 \
--cfg_scale_image=5.0 \
--cfg_scale_video=1.0 \
--vqgan_checkpoint="$vqgan_checkpoint" \
--n_frames=8 \
--mesh_dim='!1,1,-1,1' \
--dtype='fp32' \
--load_llama_config='7b' \
--update_llama_config="dict(sample_mode='vision',theta=50000000,max_sequence_length=32768,scan_attention=False,scan_query_chunk_size=128,scan_key_chunk_size=128,scan_mlp=False,scan_mlp_chunk_size=8192,scan_layers=True)" \
--load_checkpoint="$lwm_checkpoint" \
--tokenizer="$llama_tokenizer_path"
read
运行 python scripts/create_needle_data.py
目前仅支持文本和文本对话模型的 PyTorch 推理。PyTorch 模型可以作为 Hugging Face LlamaForCausalLM 模型加载。运行 python scripts/sample_pyt.py 来采样。你可能需要单独安装 torch。
有关代码库的更多详情,请参考 data.md 和 sharding.md。data.md 提供了数据处理的详情,sharding.md 提供了分片和并行度的详情。
这是基于 RingAttention 代码库,具有视觉-语言训练所需的必要特性。训练和推理已在 TPUv3 和 TPUv4 上进行了测试。
如果你遇到 bug,请开启一个 GitHub issue!
如果你使用了本代码库,或以其他方式认为我们的工作有价值,请引用:
@article{liu2023world,
title={World Model on Million-Length Video and Language with RingAttention},
author={Liu, Hao and Yan, Wilson and Zaharia, Matei and Abbeel, Pieter},
journal={arXiv preprint},
year={2024},
}
@article{liu2023ring,
title={Ring Attention with Blockwise Transformers for Near-Infinite Context},
author={Liu, Hao and Zaharia, Matei and Abbeel, Pieter},
journal={International Conference on Learning Representations},
year={2024}
}
@article{liu2023blockwise,
title={Blockwise Parallel Transformer for Large Context Models},
author={Liu, Hao and Abbeel, Pieter},
journal={Advances in neural information processing systems},
year={2023}
}
LWM 的代码在 Apache 2.0 许可证下发布。详见 LICENSE。模型在 Llama-2 许可证下发布。