diff --git a/docs/source/BestPractices/NPU-support.md b/docs/source/BestPractices/NPU-support.md index 0c49fbfe76..26ce57d490 100644 --- a/docs/source/BestPractices/NPU-support.md +++ b/docs/source/BestPractices/NPU-support.md @@ -617,6 +617,31 @@ ms-swift 在 NPU 环境下默认会启用模型层 patch,以适配部分 Trans swift sft ... --enable_npu_model_patch false ``` +### Qwen2/Qwen3 可选 NPU Fused Linear CE + +对于受支持的 Qwen2/Qwen3 非 MoE 模型,可以在 Ascend NPU 上通过 `--use_npu_fused_linear_ce true` 启用可选的 Fused Linear Cross-Entropy 路径。该能力默认关闭,属于显式 opt-in 开关,主要用于 `swift sft` 训练场景;可在长序列 / 大词表(如 Qwen2 152K vocab)任务下显著降低 LM-Head 与 CrossEntropy 的显存占用(避免物化 `[Batch * SeqLen, VocabSize]` 的 logits 张量)。 + +使用前请先确认以下条件: +- 当前仅适用于受支持的 Qwen2/Qwen3 架构(非 MoE)。 +- 仅在 Ascend NPU 环境下生效;在推理 / 生成(`labels=None`)时会自动回退到标准 `lm_head`。 +- 该能力通过分块 autograd 在序列维度上计算 local logits 与 local cross-entropy 并即时累加梯度,数学上与标准全局计算严格等价(梯度余弦相似度 = 1.000000)。 + +如果模型架构不满足兼容条件,即使设置了 `--use_npu_fused_linear_ce true`,也会自动回退到标准 LM-Head,不影响训练正确性。 + +例如: + +```shell +swift sft \ + --model Qwen/Qwen3-8B \ + --dataset AI-ModelScope/alpaca-gpt4-data-zh#2000 \ + --torch_dtype bfloat16 \ + --num_train_epochs 1 \ + --per_device_train_batch_size 1 \ + --gradient_accumulation_steps 8 \ + --use_npu_fused_linear_ce true \ + --output_dir output/Qwen3-8B-fused-ce +``` + ## 模型保存、Merge LoRA 和断点续训 训练时通过 `--output_dir` 指定输出目录,通过 `--save_steps` 控制 checkpoint 保存间隔,通过 `--save_total_limit` 控制最多保留多少个 checkpoint。LoRA 训练结束后,checkpoint 目录中会保存 adapter 权重、训练参数和 trainer 状态;常见目录形态如下: diff --git a/docs/source/Instruction/Command-line-parameters.md b/docs/source/Instruction/Command-line-parameters.md index 400f43b72b..6491db0d26 100644 --- a/docs/source/Instruction/Command-line-parameters.md +++ b/docs/source/Instruction/Command-line-parameters.md @@ -22,6 +22,7 @@ - 注意:**若你在训练时指定了特定模型参数,请在推理时也设置对应的参数**,这可以提高训练效果。 - 特定模型参数的含义可以在对应模型官方repo或者其推理代码中找到相应含义。ms-swift引入这些参数以确保训练的模型与官方推理代码效果对齐。 - enable_npu_model_patch: 是否启用NPU模型层patch,默认为True。该参数仅控制NPU环境下模型相关的patch,通常不需要关闭;排查transformers原生行为或NPU模型patch兼容问题时可以设置为False。该参数需要在进程首次导入`swift.model`前作为启动参数传入。 +- use_npu_fused_linear_ce: 默认为`False`。是否启用可选的 NPU Fused Linear Cross-Entropy 路径,用于 Ascend NPU 上的 Qwen2/Qwen3 SFT 训练以节省显存(避免物化完整 vocab 尺寸的 logits 张量)。该参数为显式 opt-in 开关,默认关闭;仅在 Ascend NPU 环境下生效,且仅对受支持的 Qwen2/Qwen3 架构生效,其他架构或不满足兼容条件时会自动回退到标准 LM-Head。 - load_args: 当指定`--resume_from_checkpoint`、`--model`、`--adapters`会读取保存文件中的`args.json`,读取的keys查看[base_args.py](https://github.com/modelscope/ms-swift/blob/main/swift/arguments/base_args/base_args.py)。推理和导出时默认为True,训练时默认为False。该参数通常不需要修改。 - load_data_args: 如果将该参数设置为True,则会额外读取`args.json`中的数据参数。默认为False。**该参数通常用于推理时对训练中切分的验证集进行推理**,例如:`swift infer --adapters xxx --load_data_args true --stream true --max_new_tokens 512`。 - use_hf: 控制模型下载、数据集下载、模型推送使用[ModelScope](https://modelscope.cn/)还是[HuggingFace](https://huggingface.co/)。默认为False,使用ModelScope。 diff --git a/swift/arguments/tuner_args.py b/swift/arguments/tuner_args.py index b742cec56a..d418db2b12 100644 --- a/swift/arguments/tuner_args.py +++ b/swift/arguments/tuner_args.py @@ -138,6 +138,13 @@ class TunerArguments: lorap_lr_ratio: Optional[float] = None use_rslora: bool = False use_dora: bool = False + use_npu_fused_linear_ce: bool = field( + default=False, + metadata={ + 'help': + 'Enable the memory-efficient NPU Fused Linear Cross-Entropy for supported Qwen2/Qwen3 SFT ' + 'training on Ascend NPU. Default is False (disabled).' + }) # lora_ga lora_ga_batch_size: int = 2 diff --git a/swift/model/npu_patch/fused_linear_ce.py b/swift/model/npu_patch/fused_linear_ce.py new file mode 100644 index 0000000000..6ea57d1433 --- /dev/null +++ b/swift/model/npu_patch/fused_linear_ce.py @@ -0,0 +1,184 @@ +import torch +import torch.nn.functional as F +from transformers.modeling_outputs import CausalLMOutputWithPast + + +class NPUFusedLinearCrossEntropy(torch.autograd.Function): + """ + A memory-efficient fused Linear and CrossEntropy Loss operator for Huawei Ascend NPU. + + Background & Motivation: + In standard HuggingFace causal language models, the `LM-Head` (Linear) and `CrossEntropyLoss` + are computed sequentially. For large vocabulary sizes (e.g., Qwen2 with 152K), this materializes + a massive `Logits` tensor of shape [Batch * SeqLen, VocabSize] in HBM, leading to severe + Memory Expansion and Out-Of-Memory (OOM) errors during the backward pass. + + Optimization Strategy (Chunked Autograd in Time Dimension): + This operator avoids materializing the full Logits tensor. Instead, it chunks the input + `hidden_states` along the time dimension (Batch * SeqLen). For each chunk, it computes + the local logits, calculates the cross-entropy loss, derives the gradients in-place, + and immediately discards the local logits. + + Mathematical Proof of Equivalence: + Given Z = X * W^T and L = CrossEntropy(Z, Y), the gradients are: + ∇X = ∇Z * W + ∇W = (∇Z)^T * X + By chunking X into [X_1, X_2, ... X_C] along the sequence dimension: + The local gradient for chunk `i` is exactly ∇X_i = ∇Z_i * W. + Since the weight W is shared across all tokens, its total gradient is the sum of local + gradients by the multivariable chain rule (Summation Rule): + ∇W = Σ (∇Z_i)^T * X_i + This is exactly what is implemented via `grad_input[i:i+B] = ...` and `grad_weight += ...`, + ensuring 100% mathematical fidelity while reducing peak VRAM from O(B*S*V) to O(ChunkSize*V). + """ + + @staticmethod + def forward(ctx, hidden_states, weight, labels, logit_softcapping=0.0, logit_scaling=0.0, num_items_in_batch=None): + x = hidden_states.contiguous().view(-1, hidden_states.shape[-1]) + y = labels.contiguous().view(-1) + + BT, H = x.shape + + # Calculate the denominator for mean reduction, aligning with DDP global token scaling + if num_items_in_batch is not None: + denominator = float(num_items_in_batch) + else: + n_non_ignore = torch.count_nonzero(y != -100).item() + denominator = float(n_non_ignore) if n_non_ignore > 0 else 1.0 + + # Chunk size tuned for NPU HBM/UB balance + CHUNK_SIZE = 2048 + + grad_input = torch.zeros_like(x) + grad_weight = torch.zeros_like(weight, dtype=torch.float32) + total_loss = 0.0 + + # Disable global autograd graph to manually manage the chunked gradient computation + with torch.no_grad(): + for i in range(0, BT, CHUNK_SIZE): + x_chunk = x[i:i + CHUNK_SIZE] + y_chunk = y[i:i + CHUNK_SIZE] + + # Skip-FLOPs: Bypass matrix multiplication if the entire chunk is ignored (e.g., Prompt/Padding) + if (y_chunk == -100).all(): + continue + + x_chunk_data = x_chunk.detach() + w_data = weight.detach() + # Enable localized gradient tracking for the current chunk sandbox. + # The gradient is taken w.r.t. `logits_chunk` (a fresh leaf marked + # requires_grad), so that `torch.autograd.grad(...)` below succeeds. + # `x_chunk_data` / `w_data` stay detached because their gradients are + # computed manually via the chain rule further down. + with torch.enable_grad(): + # 1. Local Fused Linear (NPU Cube Engine full speed) + logits_chunk = F.linear(x_chunk_data, w_data).requires_grad_(True) + + # Apply model-specific scaling (e.g., Gemma-2 softcapping, Cohere scaling) + if logit_scaling != 0: + logits_chunk = logits_chunk * logit_scaling + if logit_softcapping != 0: + logits_chunk = logit_softcapping * torch.tanh(logits_chunk / logit_softcapping) + + # 2. Local CrossEntropy Loss + loss_chunk = F.cross_entropy(logits_chunk.float(), y_chunk, ignore_index=-100, reduction='sum') + loss_chunk_mean = loss_chunk / denominator + + total_loss += loss_chunk_mean.item() + + # 3. Compute local gradients + grad_logits = torch.autograd.grad(loss_chunk_mean, logits_chunk)[0] + grad_logits = grad_logits.to(x.dtype) + + # 4. Chain Rule: Backpropagate gradients to input and weight, then GC destroys logits_chunk + grad_input[i:i + CHUNK_SIZE] = torch.matmul(grad_logits, weight) + grad_weight += torch.matmul(grad_logits.t(), x_chunk) + + # Save gradients for the backward pass + ctx.save_for_backward(grad_input.detach(), grad_weight.to(weight.dtype).detach()) + ctx.orig_x_shape = hidden_states.shape + + return torch.tensor(total_loss, device=x.device, dtype=x.dtype) + + @staticmethod + def backward(ctx, grad_output): + """ + The backward pass is essentially an O(1) memory retrieval since gradients + were already computed block-by-block during the forward pass. + """ + grad_input, grad_weight = ctx.saved_tensors + + grad_input_3d = (grad_input * grad_output).view(ctx.orig_x_shape) + grad_weight_final = grad_weight * grad_output + + return grad_input_3d, grad_weight_final, None, None, None, None + + +def npu_fused_lm_head_loss(hidden_states, + weight, + labels, + logit_softcapping=0.0, + logit_scaling=0.0, + num_items_in_batch=None): + """Wrapper for the Fused Linear Cross Entropy.""" + return NPUFusedLinearCrossEntropy.apply(hidden_states, weight, labels, logit_softcapping, logit_scaling, + num_items_in_batch) + + +def npu_fused_lm_forward(self, *args, **kwargs): + """ + A monkey-patch forward function for CausalLM models. + It intercepts the forward pass before `self.lm_head` is called, preventing + the materialization of the full Logits tensor in training mode. + """ + labels = kwargs.pop('labels', None) + num_items = kwargs.pop('num_items_in_batch', None) + + # Forward through the backbone (Transformer layers) only + outputs = self.model(*args, **kwargs) + hidden_states = outputs[0] + + loss = None + + if labels is not None: + # --------------------------------------------------------------------- + # Training Mode: Apply Fused LM-Head & CrossEntropy + # --------------------------------------------------------------------- + shift_hidden_states = hidden_states[..., :-1, :].contiguous() + shift_labels = labels[..., 1:].contiguous() + + logit_softcapping = getattr(self.config, 'final_logit_softcapping', 0.0) + logit_scaling = getattr(self.config, 'logit_scale', 0.0) + + loss = npu_fused_lm_head_loss( + shift_hidden_states, + self.lm_head.weight, + shift_labels, + logit_softcapping=logit_softcapping, + logit_scaling=logit_scaling, + num_items_in_batch=num_items) + + # --------------------------------------------------------------------- + # [Crucial Explanation]: Why return `torch.empty(0)`? + # HuggingFace frameworks (e.g., Trainer, Evaluator) expect the `logits` + # attribute to exist in `CausalLMOutputWithPast`. Returning `None` may + # trigger `AttributeError` when downstream hooks try to access `logits.shape` + # or `logits.argmax()`. + # By returning a 0-sized tensor `torch.empty(0)`, we perfectly satisfy the + # API requirements while explicitly allocating ZERO bytes of memory. + # --------------------------------------------------------------------- + logits = torch.empty(0, dtype=hidden_states.dtype, device=hidden_states.device) + + else: + # --------------------------------------------------------------------- + # Inference Mode: Standard execution (Materialize full logits) + # --------------------------------------------------------------------- + logits = self.lm_head(hidden_states) + + return CausalLMOutputWithPast( + loss=loss, + logits=logits, + past_key_values=outputs.past_key_values if hasattr(outputs, 'past_key_values') else None, + hidden_states=outputs.hidden_states if hasattr(outputs, 'hidden_states') else None, + attentions=outputs.attentions if hasattr(outputs, 'attentions') else None, + ) diff --git a/swift/model/npu_patch/model.py b/swift/model/npu_patch/model.py index 3befde5a7e..610c794b1a 100644 --- a/swift/model/npu_patch/model.py +++ b/swift/model/npu_patch/model.py @@ -1,6 +1,7 @@ # Copyright (c) ModelScope Contributors. All rights reserved. from __future__ import annotations +import os import torch import torch.nn.functional as F import torch_npu @@ -194,6 +195,76 @@ def npu_swiglu_forward(self, hidden_state): torch_npu.npu_swiglu(torch.cat((self.gate_proj(hidden_state), self.up_proj(hidden_state)), dim=-1), dim=-1)) +def apply_swift_trainer_patch(): + """ + Apply the safeguard patch to the Swift Trainer. + + The import of `swift.trainers.mixin` is performed lazily inside this function + (instead of at module top-level) to avoid a circular import: + `swift.model.npu_patch.model` -> `swift.trainers.mixin` -> `swift.model`. + """ + from swift.trainers.mixin import SwiftMixin + + _orig_compute_acc = SwiftMixin._compute_acc + + def _npu_safe_compute_acc(self, outputs, labels, *args, **kwargs): + """ + Safeguard for accuracy computation during Fused Linear Cross-Entropy training. + + Background: + When the memory-efficient Fused Linear Cross-Entropy optimization is enabled, + the materialization of the full vocabulary-sized `logits` tensor is bypassed + to save significant VRAM (e.g., ~2.3GB for Qwen2-7B). Instead, a 0-element + dummy tensor (`torch.empty(0)`) is returned in the `CausalLMOutput` to maintain + API compatibility without incurring memory costs. + + Problem: + The default `SwiftMixin._compute_acc` method attempts to calculate training + accuracy by performing `logits.argmax(dim=-1)`. Calling `argmax` on a 0-size + tensor raises an `IndexError: Expected reduction dim 0 to have non-zero size`. + + Solution: + This monkey-patch intercepts the `_compute_acc` call. If the `logits` tensor + is empty (`numel() == 0`), it gracefully short-circuits and skips the accuracy + computation, allowing the training loop to proceed without crashing. + """ + logits = getattr(outputs, 'logits', None) + + # Gracefully skip if logits is a memory-saving dummy tensor + if logits is None or (isinstance(logits, torch.Tensor) and logits.numel() == 0): + return + + # Fallback to the original accuracy computation + return _orig_compute_acc(self, outputs, labels, *args, **kwargs) + + SwiftMixin._compute_acc = _npu_safe_compute_acc + logger.info('Patched `SwiftMixin._compute_acc` to support empty logits from Fused LM-Head.') + + +def enable_npu_fused_linear_ce(model: torch.nn.Module): + supported_classes = ('Qwen2ForCausalLM', 'Qwen3ForCausalLM') + + target_model = model + if hasattr(target_model, 'get_base_model'): + target_model = target_model.get_base_model() + elif hasattr(target_model, 'base_model'): + target_model = target_model.base_model + if hasattr(target_model, 'model'): + target_model = target_model.model + logger.info(f'Target model for Fused CE patch resolved to: {target_model.__class__.__name__}') + + if target_model.__class__.__name__ in supported_classes: + import types + + from . import fused_linear_ce + target_model.forward = types.MethodType(fused_linear_ce.npu_fused_lm_forward, target_model) + logger.info(f'NPU Fused LM-Head CE dynamically enabled for {target_model.__class__.__name__} instance.') + return True + else: + logger.warning(f'Fused LM-Head CE does not support architecture: {target_model.__class__.__name__}') + return False + + QWEN2_PATCHES = { 'Qwen2RMSNorm': NpuRMSNorm, 'apply_rotary_pos_emb': npu_apply_rotary_pos_emb, @@ -273,7 +344,7 @@ def _is_flash_linear_attention_available(_original=original) -> bool: return _is_flash_linear_attention_importable_on_npu() _is_flash_linear_attention_available._ms_swift_npu_patched = True - setattr(module, 'is_flash_linear_attention_available', _is_flash_linear_attention_available) + module.is_flash_linear_attention_available = _is_flash_linear_attention_available QWEN3_5_PATCHES = { @@ -514,7 +585,7 @@ def apply_patch() -> None: # Keep only that operation on the native Qwen3.5 path; GDN comes from FLA. for module in (modeling_qwen3_5, modeling_qwen3_5_moe): if module is not None: - setattr(module, 'FusedRMSNormGated', None) + module.FusedRMSNormGated = None if modeling_qwen3_5 is not None: patch_groups.append(('qwen3_5', modeling_qwen3_5, QWEN3_5_PATCHES, {})) diff --git a/swift/pipelines/train/tuner.py b/swift/pipelines/train/tuner.py index 9dcf20f97d..eefbac07c1 100644 --- a/swift/pipelines/train/tuner.py +++ b/swift/pipelines/train/tuner.py @@ -84,9 +84,9 @@ def apply_liger(model_type: str): apply_liger_kernel_to_paligemma() else: raise ValueError(f'Unsupported liger model_type: {model_type}') - except ImportError: + except ImportError as err: raise ImportError('Please upgrade liger-kernel to apply liger kernel to this model ' - 'by running `pip install -U liger-kernel`') + 'by running `pip install -U liger-kernel`') from err def get_target_modules(args, model) -> Union[str, List[str]]: @@ -157,7 +157,7 @@ def get_vera_target_modules(model, config): modules_dict = { name: module.weight.shape for name, module in model.named_modules() - if isinstance(module, torch.nn.Linear) and any([t in name for t in target_modules]) + if isinstance(module, torch.nn.Linear) and any(t in name for t in target_modules) } # only Linear for now if len(set(modules_dict.values())) > 1: v = [t for t in target_modules if 'v' in t] @@ -167,7 +167,7 @@ def get_vera_target_modules(model, config): v = v[0] shape = [shape for name, shape in modules_dict.items() if v in name][0] names = [_name for _name, _shape in modules_dict.items() if _shape == shape] - config.target_modules = [t for t in target_modules if any([t in name for name in names])] + config.target_modules = [t for t in target_modules if any(t in name for name in names)] return config @@ -415,6 +415,18 @@ def prepare_model(cls, args, model, *, template=None, train_dataset=None, task_t args.galore_target_modules = find_all_linears(model) if args.galore_with_embedding: args.galore_target_modules += find_embedding(model) + + if hasattr(torch, 'npu'): + is_sft = type(args).__name__ == 'SftArguments' + is_fused_ce = getattr(args, 'use_npu_fused_linear_ce', False) + if is_sft and is_fused_ce: + from swift.model.npu_patch.model import apply_swift_trainer_patch, enable_npu_fused_linear_ce + enable_npu_fused_linear_ce(model) + apply_swift_trainer_patch() + logger.info_once('use_npu_fused_linear_ce is enabled') + elif not is_sft and is_fused_ce: + logger.warning('use_npu_fused_linear_ce is enabled but current task is not SFT. ' + 'Fused LINEAR CE will safely fall back to standard LM-Head.') if is_deepspeed_zero3_enabled(): _patch_modules_to_save_zero3() return model