执行如下命令后报错
<span class="line"><span style="color: #81A1C1">from</span><span style="color: #D8DEE9FF"> llama </span><span style="color: #81A1C1">import</span><span style="color: #D8DEE9FF"> tokenizer</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF"> Llama</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF"> Dialog</span></span>
<span class="line"><span style="color: #D8DEE9FF">checkpoint_dir </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">"</span><span style="color: #A3BE8C">/training-data/pakcages/llama/llama-2-7b-chat</span><span style="color: #ECEFF4">"</span></span>
<span class="line"><span style="color: #D8DEE9FF">tokenizer_path </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">"</span><span style="color: #A3BE8C">/training-data/pakcages/llama/tokenizer.model</span><span style="color: #ECEFF4">"</span></span>
<span class="line"><span style="color: #D8DEE9FF">temperature </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> </span><span style="color: #B48EAD">0.75</span></span>
<span class="line"><span style="color: #D8DEE9FF">top_p </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> </span><span style="color: #B48EAD">0.9</span></span>
<span class="line"><span style="color: #D8DEE9FF">max_seq_len </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> </span><span style="color: #B48EAD">128</span></span>
<span class="line"><span style="color: #D8DEE9FF">max_gen_len </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> </span><span style="color: #B48EAD">64</span></span>
<span class="line"><span style="color: #D8DEE9FF">max_batch_size </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> </span><span style="color: #B48EAD">4</span></span>
<span class="line"></span>
<span class="line"><span style="color: #D8DEE9FF">generator </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> Llama</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">build</span><span style="color: #ECEFF4">(</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">ckpt_dir</span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF">checkpoint_dir</span><span style="color: #ECEFF4">,</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">tokenizer_path</span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF">tokenizer_path</span><span style="color: #ECEFF4">,</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">max_seq_len</span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF">max_seq_len</span><span style="color: #ECEFF4">,</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">max_batch_size</span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF">max_batch_size</span><span style="color: #ECEFF4">)</span></span>ValueError: Error initializing torch.distributed using env:// rendezvous: environment variable RANK expected, but not set
是源码里面这一段引起的:
<span class="line"><span style="color: #81A1C1">if</span><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">not</span><span style="color: #D8DEE9FF"> torch</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">distributed</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">is_initialized</span><span style="color: #ECEFF4">():</span></span>
<span class="line"><span style="color: #D8DEE9FF"> torch</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">distributed</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">init_process_group</span><span style="color: #ECEFF4">(</span><span style="color: #ECEFF4">"</span><span style="color: #A3BE8C">nccl</span><span style="color: #ECEFF4">"</span><span style="color: #ECEFF4">)</span></span>启动不起来看样子是因为分布式的问题。我尝试绕开分布式,从它的build函数开始看:
<span class="line"><span style="color: #D8DEE9FF">checkpoints </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> </span><span style="color: #88C0D0">sorted</span><span style="color: #ECEFF4">(</span><span style="color: #88C0D0">Path</span><span style="color: #ECEFF4">(</span><span style="color: #D8DEE9FF">ckpt_dir</span><span style="color: #ECEFF4">).</span><span style="color: #88C0D0">glob</span><span style="color: #ECEFF4">(</span><span style="color: #ECEFF4">"</span><span style="color: #A3BE8C">*.pth</span><span style="color: #ECEFF4">"</span><span style="color: #ECEFF4">))</span></span>
<span class="line"><span style="color: #D8DEE9FF">ckpt_path </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> checkpoints</span><span style="color: #ECEFF4">[</span><span style="color: #88C0D0">get_model_parallel_rank</span><span style="color: #ECEFF4">()]</span></span>
<span class="line"><span style="color: #D8DEE9FF">checkpoint </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> torch</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">load</span><span style="color: #ECEFF4">(</span><span style="color: #D8DEE9FF">ckpt_path</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">map_location</span><span style="color: #81A1C1">=</span><span style="color: #ECEFF4">"</span><span style="color: #A3BE8C">cpu</span><span style="color: #ECEFF4">"</span><span style="color: #ECEFF4">)</span></span>检查给定的checkpoint_dir是否包含pth文件,7B的模型只有一个pth文件,所以一个进程就可以了,我想get_model_parallel_rank()大概意思即是有几个文件就启动多少个进程,代码来自facebook团队开发并行训练包fairscale:
<span class="line"><span style="color: #81A1C1">def</span><span style="color: #D8DEE9FF"> </span><span style="color: #88C0D0">get_model_parallel_rank</span><span style="color: #ECEFF4">()</span><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">-></span><span style="color: #D8DEE9FF"> </span><span style="color: #88C0D0">int</span><span style="color: #ECEFF4">:</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">"""</span><span style="color: #A3BE8C">Return my rank for the model parallel group.</span><span style="color: #ECEFF4">"""</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">return</span><span style="color: #D8DEE9FF"> torch</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">distributed</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">get_rank</span><span style="color: #ECEFF4">(</span><span style="color: #D8DEE9">group</span><span style="color: #81A1C1">=</span><span style="color: #88C0D0">get_model_parallel_group</span><span style="color: #ECEFF4">())</span></span>
<span class="line"></span>
<span class="line"><span style="color: #81A1C1">def</span><span style="color: #D8DEE9FF"> </span><span style="color: #88C0D0">get_model_parallel_group</span><span style="color: #ECEFF4">()</span><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">-></span><span style="color: #D8DEE9FF"> torch</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">distributed</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">ProcessGroup</span><span style="color: #ECEFF4">:</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">"""</span><span style="color: #A3BE8C">Get the model parallel group the caller rank belongs to.</span><span style="color: #ECEFF4">"""</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">assert</span><span style="color: #D8DEE9FF"> _MODEL_PARALLEL_GROUP </span><span style="color: #81A1C1">is</span><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">not</span><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">None</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">"</span><span style="color: #A3BE8C">model parallel group is not initialized</span><span style="color: #ECEFF4">"</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">return</span><span style="color: #D8DEE9FF"> _MODEL_PARALLEL_GROUP</span></span>咱不管这些,7B反正就一个参数文件,直接从文件夹加载:
<span class="line"><span style="color: #81A1C1">from</span><span style="color: #D8DEE9FF"> pathlib </span><span style="color: #81A1C1">import</span><span style="color: #D8DEE9FF"> Path</span></span>
<span class="line"><span style="color: #D8DEE9FF">checkpoints </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> </span><span style="color: #88C0D0">sorted</span><span style="color: #ECEFF4">(</span><span style="color: #88C0D0">Path</span><span style="color: #ECEFF4">(</span><span style="color: #D8DEE9FF">checkpoint_dir</span><span style="color: #ECEFF4">).</span><span style="color: #88C0D0">glob</span><span style="color: #ECEFF4">(</span><span style="color: #ECEFF4">"</span><span style="color: #A3BE8C">*.pth</span><span style="color: #ECEFF4">"</span><span style="color: #ECEFF4">))</span></span>
<span class="line"><span style="color: #616E88"># Llama-2-7b model weights are distributed in a single file.</span></span>
<span class="line"><span style="color: #D8DEE9FF">checkpoint </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> checkpoints</span><span style="color: #ECEFF4">[</span><span style="color: #B48EAD">0</span><span style="color: #ECEFF4">]</span></span>
<span class="line"><span style="color: #D8DEE9FF">device </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> torch</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">device</span><span style="color: #ECEFF4">(</span><span style="color: #ECEFF4">'</span><span style="color: #A3BE8C">cuda</span><span style="color: #ECEFF4">'</span><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">if</span><span style="color: #D8DEE9FF"> torch</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">cuda</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">is_available</span><span style="color: #ECEFF4">()</span><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">else</span><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">'</span><span style="color: #A3BE8C">cpu</span><span style="color: #ECEFF4">'</span><span style="color: #ECEFF4">)</span></span>
<span class="line"><span style="color: #D8DEE9FF">checkpoint </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> torch</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">load</span><span style="color: #ECEFF4">(</span><span style="color: #D8DEE9FF">checkpoint</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">map_location</span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF">device</span><span style="color: #ECEFF4">)</span></span>7b模型直接加载进显卡占用了13671MB的显存:

huggingface转换
算了,还是先不重写了,先用huggingface转换吧,安装一下transformers,他的转换函数在src/transformers/models/llama/convert_llama_weights_to_hf.py 这里可以看源文件。
本身安装transformers的时候已经安装了这个模块,写个脚本:
<span class="line"><span style="color: #81A1C1">from</span><span style="color: #D8DEE9FF"> transformers</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">models</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">llama</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">convert_llama_weights_to_hf </span><span style="color: #81A1C1">import</span><span style="color: #D8DEE9FF"> main</span></span>
<span class="line"></span>
<span class="line"><span style="color: #81A1C1">if</span><span style="color: #D8DEE9FF"> __name__ </span><span style="color: #81A1C1">==</span><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">"</span><span style="color: #A3BE8C">__main__</span><span style="color: #ECEFF4">"</span><span style="color: #ECEFF4">:</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #88C0D0">main</span><span style="color: #ECEFF4">()</span></span>直接这个脚本就行了,执行一下:
<span class="line"><span style="color: #616E88"># python convert.py --help</span></span>
<span class="line"><span style="color: #D8DEE9FF">usage</span><span style="color: #ECEFF4">:</span><span style="color: #D8DEE9FF"> convert</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">py </span><span style="color: #ECEFF4">[</span><span style="color: #81A1C1">-</span><span style="color: #D8DEE9FF">h</span><span style="color: #ECEFF4">]</span><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">[</span><span style="color: #D8DEE9">--</span><span style="color: #D8DEE9FF">input_dir INPUT_DIR</span><span style="color: #ECEFF4">]</span><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">[</span><span style="color: #D8DEE9">--</span><span style="color: #D8DEE9FF">model_size </span><span style="color: #ECEFF4">{</span><span style="color: #D8DEE9">7B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">7Bf</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">13B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">13Bf</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">30B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">34B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">65B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">70B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">70Bf</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF">tokenizer_only</span><span style="color: #ECEFF4">}]</span><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">[</span><span style="color: #D8DEE9">--</span><span style="color: #D8DEE9FF">output_dir OUTPUT_DIR</span><span style="color: #ECEFF4">]</span><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">[</span><span style="color: #D8DEE9">--</span><span style="color: #D8DEE9FF">safe_serialization SAFE_SERIALIZATION</span><span style="color: #ECEFF4">]</span></span>
<span class="line"></span>
<span class="line"><span style="color: #D8DEE9FF">options</span><span style="color: #ECEFF4">:</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">-</span><span style="color: #D8DEE9FF">h</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">--</span><span style="color: #88C0D0">help</span><span style="color: #D8DEE9FF"> show this </span><span style="color: #88C0D0">help</span><span style="color: #D8DEE9FF"> message </span><span style="color: #81A1C1">and</span><span style="color: #D8DEE9FF"> </span><span style="color: #88C0D0">exit</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">--</span><span style="color: #D8DEE9FF">input_dir INPUT_DIR</span></span>
<span class="line"><span style="color: #D8DEE9FF"> Location of LLaMA weights</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF"> which contains tokenizer</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">model </span><span style="color: #81A1C1">and</span><span style="color: #D8DEE9FF"> model folders</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">--</span><span style="color: #D8DEE9FF">model_size </span><span style="color: #ECEFF4">{</span><span style="color: #D8DEE9">7B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">7Bf</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">13B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">13Bf</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">30B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">34B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">65B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">70B</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9">70Bf</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF">tokenizer_only</span><span style="color: #ECEFF4">}</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #ECEFF4">'</span><span style="color: #A3BE8C">f</span><span style="color: #ECEFF4">'</span><span style="color: #D8DEE9FF"> models correspond to the finetuned versions</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">and</span><span style="color: #D8DEE9FF"> are specific to the Llama2 official release</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF"> For more details on Llama2</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF"> checkout the original repo</span><span style="color: #ECEFF4">:</span><span style="color: #D8DEE9FF"> https</span><span style="color: #ECEFF4">:</span><span style="color: #81A1C1">//</span><span style="color: #D8DEE9FF">huggingface</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">co</span><span style="color: #81A1C1">/</span><span style="color: #D8DEE9FF">meta</span><span style="color: #81A1C1">-</span><span style="color: #D8DEE9FF">llama</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">--</span><span style="color: #D8DEE9FF">output_dir OUTPUT_DIR</span></span>
<span class="line"><span style="color: #D8DEE9FF"> Location to write HF model </span><span style="color: #81A1C1">and</span><span style="color: #D8DEE9FF"> tokenizer</span></span>
<span class="line"><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">--</span><span style="color: #D8DEE9FF">safe_serialization SAFE_SERIALIZATION</span></span>
<span class="line"><span style="color: #D8DEE9FF"> Whether </span><span style="color: #81A1C1">or</span><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">not</span><span style="color: #D8DEE9FF"> to save using </span><span style="color: #D8DEE9">`safetensors`</span><span style="color: #ECEFF4">.</span></span>–input_dir 写的是llama的根目录
–model_size 选择要转换的模型参数量,这个有个bug,你只能填提供的那几个名字,问题是llama的目录下对应的模型文件名是”llama-2-*b”这种,转换脚本会去”input_dir/*B”下面去找模型文件,所以需要给”llama-2-*b”改成”*B”后再执行脚本。
转换过程中并不需要GPU,完成之后用transformers加载就行了,7B的模型加载完后划分了26G现存

执行:
<span class="line"><span style="color: #D8DEE9FF">total_params </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> </span><span style="color: #88C0D0">sum</span><span style="color: #ECEFF4">(</span><span style="color: #D8DEE9FF">p</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">numel</span><span style="color: #ECEFF4">()</span><span style="color: #D8DEE9FF"> </span><span style="color: #81A1C1">for</span><span style="color: #D8DEE9FF"> p </span><span style="color: #81A1C1">in</span><span style="color: #D8DEE9FF"> model</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">parameters</span><span style="color: #ECEFF4">())</span></span>
<span class="line"><span style="color: #88C0D0">print</span><span style="color: #ECEFF4">(</span><span style="color: #81A1C1">f</span><span style="color: #A3BE8C">"Total number of parameters: </span><span style="color: #EBCB8B">{</span><span style="color: #D8DEE9FF">total_params</span><span style="color: #EBCB8B">}</span><span style="color: #A3BE8C">"</span><span style="color: #ECEFF4">)</span></span>
<span class="line"><span style="color: #81A1C1">>>></span><span style="color: #D8DEE9FF"> Total number of parameters</span><span style="color: #ECEFF4">:</span><span style="color: #D8DEE9FF"> </span><span style="color: #B48EAD">6607343616</span></span>模型的总参数是6,607,343,616
粗略计算一下,采用单精度(single-precision float-point format)存储这些参数的话总共要用6607343616 * 32 / 8 / 1024 / 1024 / 1024 = 24.6143G,如果采用半精度存储看下:
<span class="line"><span style="color: #D8DEE9FF">tokenizer </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> LlamaTokenizer</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">from_pretrained</span><span style="color: #ECEFF4">(</span><span style="color: #ECEFF4">'</span><span style="color: #A3BE8C">llama_hf/7Bf</span><span style="color: #ECEFF4">'</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">torch_dtype</span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF">torch</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">float16</span><span style="color: #ECEFF4">)</span></span>
<span class="line"><span style="color: #D8DEE9FF">model </span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF"> LlamaModel</span><span style="color: #ECEFF4">.</span><span style="color: #88C0D0">from_pretrained</span><span style="color: #ECEFF4">(</span><span style="color: #ECEFF4">'</span><span style="color: #A3BE8C">llama_hf/7Bf</span><span style="color: #ECEFF4">'</span><span style="color: #ECEFF4">,</span><span style="color: #D8DEE9FF"> </span><span style="color: #D8DEE9">torch_dtype</span><span style="color: #81A1C1">=</span><span style="color: #D8DEE9FF">torch</span><span style="color: #ECEFF4">.</span><span style="color: #D8DEE9FF">float16</span><span style="color: #ECEFF4">)</span></span>放进显卡的话划分了13G现存

还想再降低显存占用就得使用quantization了,后面整理下。
LlamaForCausalLM可以用来生成回答,默认的LlamaModel没有这功能