01 背景和目标
- **目标:**在 Ascend 上用 MindSpore 跑通 Llama(推理 + 微调),尽量少魔改,支持 KV Cache、RoPE、混合精度和断点恢复。
- **限制:**不依赖奇怪分支;只用公开可得的接口(MindSpore 基座 + 常见组件)。
- **策略:**能复用的就复用(Tokenizer、权重),不能复用的就写一个薄转换层。不追求一步到位,但要'能打'。
02 环境要点
MindSpore 有两种模式:GRAPH_MODE(编译图)和 PYNATIVE_MODE(动态图)。在 Ascend 上尽量用 GRAPH,性能差一大截不是开玩笑的。
import mindspore as ms
ms.set_context(mode=ms.GRAPH_MODE, device_target="Ascend")
# 可选:减少首次编译抖动
ms.set_context(jit_config={"jit_level": "O2"}) # 视版本而定
混合精度推荐 O2,配合 loss scale(训练阶段):
from mindspore.amp import auto_mixed_precision, StaticLossScaler
net = build_llama() # 你自己的 Llama Cell
auto_mixed_precision(net, "O2") # 权重/计算多落到 fp16/bf16
loss_scaler = StaticLossScaler(2**12)
⚠️ 注意:MindSpore 对 Ascend 的算子融合比较激进,图模式下某些自定义 Python 控制流容易被'优化没了'。遇到莫名其妙的数值波动,先关掉你新加的'聪明'控制流。
03 Tokenizer 与 RoPE:别在细节上翻车
- **Tokenizer:**我直接复用 HF 的 tokenizer.json 和 tokenizer.model,在数据前处理阶段完成编码解码。训练/推理时只给 MindSpore 喂 input_ids 和 attention_mask(注意 mask 的 dtype 和 shape)。
- **RoPE(Rotary Embedding):**MindSpore 里实现 RoPE 时,位置索引的广播维度和角度表(cos/sin)缓存要提前考虑到 prefill+decode 两阶段。 简化做法:预缓存最大 max_seq_len 的 cos/sin;decode 阶段按 pos_offset 索引切片。
def precompute_rope(theta_base, head_dim, max_len, dtype=ms.float16):
inv_freq = 1.0 / (theta_base ** (ms.numpy.arange(0, head_dim, 2, dtype=ms.float32) / head_dim))
t = ms.numpy.arange(max_len, dtype=ms.float32)
freqs = ms.numpy.einsum('n,d->nd', t, inv_freq)
cos = ms.numpy.cos(freqs).astype(dtype)
sin = ms.numpy.sin(freqs).astype(dtype)
cos, sin
():
cos_t = cos[pos]
sin_t = sin[pos]
_ ():
cos_t = ms.ops.expand_dims(cos_t, )
sin_t = ms.ops.expand_dims(sin_t, )
cos_t = ms.ops.expand_dims(cos_t, )
sin_t = ms.ops.expand_dims(sin_t, )
q1, q2 = q[..., ::], q[..., ::]
k1, k2 = k[..., ::], k[..., ::]
q_rot = ms.ops.stack([q1 * cos_t - q2 * sin_t, q1 * sin_t + q2 * cos_t], axis=-).reshape(q.shape)
k_rot = ms.ops.stack([k1 * cos_t - k2 * sin_t, k2 * sin_t + k1 * cos_t], axis=-).reshape(k.shape)
q_rot, k_rot

