Llama中模块参数大小

LLama2中,流程中数据大小的变换如下

Transformer模块

第一次输入,进行prefill,输入x维度为1, 8, 4096

  1. 构建wq,wk,wv,wo,尺寸均为4096,4096, 与x点乘,得到xq, xk, xv

  2. 构建KV cache, 尺寸为 batch size, max_seq_len, local_kv_heads, head_dim,对应 1, 8, 32, 128

3.基于kv cache构造 keys, alues,对应的尺寸还是1,8,32,128

  1. 在最后两个维度对于xq和key进行点乘,得到scores,维度变成【1, 32, 8, 8】

  2. 将mask与scores相加

  3. 对于scores进行softmax

  4. 将scores 1, 32, 8, 8与values 1, 32, 8, 128进行乘法

  5. 得到output 1, 8, 4096

  6. 将output再与wo进行乘法1, 8, 4096

  7. 接下来对于输出进行 ffn_norm的操作

Feedforward模块

11.然后进行feed_forward.得到当前transformer模块的输出 1, 8, 4096

feed_forward的操作如下,虽然代码很小,但是计算量却很大。

复制代码
    def forward(self, x):
        return self.w2(F.silu(self.w1(x)) * self.w3(x))

其中,w1的维度为11008, 4096, w2的维度为4096, 11008, w3的维度为11008, 4096

kv cache的表达如下

python 复制代码
        self.cache_k = torch.zeros(
            (
                args.max_batch_size,
                args.max_seq_len,
                self.n_local_kv_heads,
                self.head_dim,
            )
        ).cuda()
        self.cache_v = torch.zeros(
            (
                args.max_batch_size,
                args.max_seq_len,
                self.n_local_kv_heads,
                self.head_dim,
            )

关于kv cache的细节讨论

llama2设定 local_kv_heads为32,head_dim为128。所以,kv cache的尺寸为 1, 512,32, 128 * 2

对于一个batch的数据来说哦,因为llama2 7B 包含32个transformer,所以,当使用FP32表达时, 对应一个batch的kv cache的大小为128 * 32 * 128 *2 * 32 * 4byte= 0.5GB.

这里,也可以看到几个变量:

* 当batch变大时,kv cache线性增长

* 当batch 的最大长度增大时, Kv cache线性增长。

参考链接:

https://arxiv.org/pdf/1911.02150

相关推荐
图王大胜4 分钟前
万物演化论07(第三章 ) 文明启动,开始有了意识
人工智能·ai·宇宙·演化·系统科学
QC777LX7 分钟前
制造业成本会计学怎么学AI,从哪几个工作环节开始更好?
人工智能
DolphinDB智臾科技7 分钟前
不懂复杂金融数据,也能让 AI 做投研:DolphinDB 股票分析 MCP 已开源
人工智能·金融·开源
霸道流氓气质9 分钟前
Spring AI 技术细节:VectorStore 多库统一抽象
java·人工智能·spring
AI天行健9 分钟前
文生视频与图生视频的技术区别及适用场景分析
人工智能·音视频
今天AI了吗10 分钟前
Codex 配置自定义 AI API 完整指南:从零到一接入你的专属模型
java·人工智能·python·数据分析·embedding
东离与糖宝10 分钟前
不用高端显卡!本地大模型量化入门|Ollama+transformers+llama.cpp实战
人工智能
森山冶仁10 分钟前
治理知识库构建:用 RAG 把制度、文档、经验变成 AI 能力
人工智能·智能问答·rag·企业知识库·大模型落地·ai治理
安科瑞黄益鸣17 分钟前
筑牢配电安全:安科瑞 ARB 弧光保护在半导体厂房的应用
人工智能
程序猿编码17 分钟前
基于GGML的C++17轻量化语音推理引擎:说话人识别与语音分析技术全解析
开发语言·c++·pytorch·深度学习·神经网络·大模型