mirror of
https://github.com/hiyouga/LLaMA-Factory.git
synced 2026-08-22 07:55:46 +08:00
[train] Add KTransformers VLM fine-tuning support (#10760)
This commit is contained in:
committed by
GitHub
parent
f28afaf635
commit
ff6d4d12ee
@@ -0,0 +1,51 @@
|
|||||||
|
### model
|
||||||
|
model_name_or_path: Qwen/Qwen3-VL-30B-A3B-Instruct
|
||||||
|
image_max_pixels: 262144
|
||||||
|
video_max_pixels: 16384
|
||||||
|
trust_remote_code: true
|
||||||
|
|
||||||
|
### method
|
||||||
|
stage: sft
|
||||||
|
do_train: true
|
||||||
|
finetuning_type: lora
|
||||||
|
lora_rank: 8
|
||||||
|
lora_alpha: 16
|
||||||
|
lora_target: all
|
||||||
|
|
||||||
|
### dataset
|
||||||
|
dataset: mllm_demo
|
||||||
|
template: qwen3_vl
|
||||||
|
cutoff_len: 512
|
||||||
|
overwrite_cache: true
|
||||||
|
preprocessing_num_workers: 4
|
||||||
|
dataloader_num_workers: 1
|
||||||
|
|
||||||
|
### output
|
||||||
|
output_dir: saves/KT_FT_qwen3vl_30b_a3b
|
||||||
|
logging_steps: 1
|
||||||
|
save_steps: 100
|
||||||
|
plot_loss: true
|
||||||
|
overwrite_output_dir: true
|
||||||
|
report_to: none
|
||||||
|
|
||||||
|
### train
|
||||||
|
per_device_train_batch_size: 1
|
||||||
|
gradient_accumulation_steps: 1
|
||||||
|
learning_rate: 1.0e-4
|
||||||
|
num_train_epochs: 1.0
|
||||||
|
lr_scheduler_type: cosine
|
||||||
|
warmup_ratio: 0.1
|
||||||
|
bf16: true
|
||||||
|
ddp_timeout: 360000000
|
||||||
|
|
||||||
|
### ktransformers
|
||||||
|
use_kt: true
|
||||||
|
kt_cpu_activation: retain
|
||||||
|
kt_config:
|
||||||
|
kt_expert_weight_format: bf16
|
||||||
|
kt_backend: AMXBF16
|
||||||
|
kt_num_threads: 96
|
||||||
|
kt_tp_enabled: true
|
||||||
|
kt_threadpool_count: 2
|
||||||
|
kt_max_cache_depth: 2
|
||||||
|
kt_share_backward_bb: true
|
||||||
@@ -558,6 +558,12 @@ class KTransformersArguments:
|
|||||||
raise TypeError("LLaMA-Factory `kt_config` must be a flat mapping.")
|
raise TypeError("LLaMA-Factory `kt_config` must be a flat mapping.")
|
||||||
|
|
||||||
config = dict(raw_config)
|
config = dict(raw_config)
|
||||||
|
# Transformers-KT adds these defaults to the user-owned mapping during
|
||||||
|
# TrainingArguments.__post_init__. They are transport metadata, not
|
||||||
|
# conflicting user overrides.
|
||||||
|
for key in ("enabled", "kt_skip_expert_loading"):
|
||||||
|
if config.get(key) is True:
|
||||||
|
config.pop(key)
|
||||||
conflicts = sorted(set(config) & self._KT_DERIVED_KEYS)
|
conflicts = sorted(set(config) & self._KT_DERIVED_KEYS)
|
||||||
if conflicts:
|
if conflicts:
|
||||||
raise ValueError(f"These `kt_config` values are derived from LLaMA-Factory arguments: {conflicts}.")
|
raise ValueError(f"These `kt_config` values are derived from LLaMA-Factory arguments: {conflicts}.")
|
||||||
@@ -608,8 +614,9 @@ class KTransformersArguments:
|
|||||||
advanced_config: dict[str, Any] | None = None,
|
advanced_config: dict[str, Any] | None = None,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
r"""Map LLaMA-Factory-owned training values to the public KT configuration."""
|
r"""Map LLaMA-Factory-owned training values to the public KT configuration."""
|
||||||
if getattr(finetuning_args, "finetuning_type", None) != "lora":
|
finetuning_type = getattr(finetuning_args, "finetuning_type", None)
|
||||||
raise ValueError("KTransformers thin integration currently supports LoRA finetuning only.")
|
if finetuning_type not in {"lora", "full"}:
|
||||||
|
raise ValueError("KTransformers supports LoRA and full-parameter finetuning.")
|
||||||
|
|
||||||
kt_config = dict(advanced_config or {})
|
kt_config = dict(advanced_config or {})
|
||||||
configured_capacity = kt_config.pop("kt_model_max_length", None)
|
configured_capacity = kt_config.pop("kt_model_max_length", None)
|
||||||
@@ -634,8 +641,8 @@ class KTransformersArguments:
|
|||||||
"kt_lora_expert_num": self.kt_lora_expert_num,
|
"kt_lora_expert_num": self.kt_lora_expert_num,
|
||||||
"kt_lora_expert_intermediate_size": self.kt_lora_expert_intermediate_size,
|
"kt_lora_expert_intermediate_size": self.kt_lora_expert_intermediate_size,
|
||||||
"kt_activation_policy": self.get_kt_activation_policy(),
|
"kt_activation_policy": self.get_kt_activation_policy(),
|
||||||
"kt_train_mode": "lora",
|
"kt_train_mode": finetuning_type,
|
||||||
"kt_full_weight_grad": False,
|
"kt_full_weight_grad": finetuning_type == "full",
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
return {key: value for key, value in kt_config.items() if value is not None}
|
return {key: value for key, value in kt_config.items() if value is not None}
|
||||||
@@ -668,13 +675,29 @@ class KTransformersArguments:
|
|||||||
self._get_advanced_kt_config(training_args),
|
self._get_advanced_kt_config(training_args),
|
||||||
)
|
)
|
||||||
update_kt_config = getattr(training_args, "update_kt_config", None)
|
update_kt_config = getattr(training_args, "update_kt_config", None)
|
||||||
if not callable(update_kt_config):
|
|
||||||
raise RuntimeError(
|
|
||||||
"The installed Transformers-KT does not provide `TrainingArguments.update_kt_config()`."
|
|
||||||
)
|
|
||||||
|
|
||||||
adapter_dir = self._resolve_kt_adapter_artifact_dir("training")
|
adapter_dir = self._resolve_kt_adapter_artifact_dir("training")
|
||||||
update_kt_config(kt_config, adapter_name_or_path=adapter_dir)
|
if callable(update_kt_config):
|
||||||
|
update_kt_config(kt_config, adapter_name_or_path=adapter_dir)
|
||||||
|
return
|
||||||
|
|
||||||
|
# transformers-kt 5.6 exposes the config object but predates the public
|
||||||
|
# update helper. Keep its flat loading config and Accelerate's nested
|
||||||
|
# plugin config synchronized without requiring another dependency pin.
|
||||||
|
from kt_kernel.sft import KTConfig
|
||||||
|
|
||||||
|
supported_keys = {item.name for item in fields(KTConfig)}
|
||||||
|
compatible_config = {key: value for key, value in kt_config.items() if key in supported_keys}
|
||||||
|
if self.get_kt_activation_policy()["gpu"] == "recompute":
|
||||||
|
compatible_config.setdefault("kt_share_cache_pool", True)
|
||||||
|
|
||||||
|
hf_kt_config = getattr(training_args, "hf_kt_config", None)
|
||||||
|
if hf_kt_config is None or not isinstance(getattr(hf_kt_config, "_kt_config", None), dict):
|
||||||
|
raise RuntimeError("The installed Transformers-KT does not expose a mutable KT configuration.")
|
||||||
|
hf_kt_config._kt_config.update(compatible_config)
|
||||||
|
|
||||||
|
accelerator_config = getattr(training_args, "accelerator_config", None)
|
||||||
|
if accelerator_config is not None:
|
||||||
|
accelerator_config.kt_config = {"enabled": True, "kt_config": compatible_config}
|
||||||
|
|
||||||
def configure_kt_loading(self, finetuning_args: Any, model_max_length: int | None) -> None:
|
def configure_kt_loading(self, finetuning_args: Any, model_max_length: int | None) -> None:
|
||||||
r"""Configure KT model loading for inference and evaluation."""
|
r"""Configure KT model loading for inference and evaluation."""
|
||||||
|
|||||||
@@ -196,13 +196,22 @@ def load_model(
|
|||||||
|
|
||||||
# Conv3D is not recommended when using torch 2.9.x
|
# Conv3D is not recommended when using torch 2.9.x
|
||||||
if is_torch_version_greater_than("2.9.0") and not is_torch_version_greater_than("2.10.0"):
|
if is_torch_version_greater_than("2.9.0") and not is_torch_version_greater_than("2.10.0"):
|
||||||
if any(isinstance(m, torch.nn.Conv3d) for m in model.modules()):
|
conv3d_modules = [module for module in model.modules() if isinstance(module, torch.nn.Conv3d)]
|
||||||
|
kt_conv3d_ready = (
|
||||||
|
model_args.use_kt
|
||||||
|
and is_trainable
|
||||||
|
and bool(conv3d_modules)
|
||||||
|
and all(getattr(module, "_kt_conv3d_compatible", False) for module in conv3d_modules)
|
||||||
|
)
|
||||||
|
if conv3d_modules and not kt_conv3d_ready:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Unsupported torch version detected: torch 2.9.x with Conv3D. "
|
"Unsupported torch version detected: torch 2.9.x with Conv3D. "
|
||||||
"This combination is known to cause severe performance regression. "
|
"This combination is known to cause severe performance regression. "
|
||||||
"Please downgrade torch to <2.9 or remove Conv3D. "
|
"Please downgrade torch to <2.9 or remove Conv3D. "
|
||||||
"See https://github.com/pytorch/pytorch/issues/166122"
|
"See https://github.com/pytorch/pytorch/issues/166122"
|
||||||
)
|
)
|
||||||
|
elif kt_conv3d_ready:
|
||||||
|
logger.info_rank0("Using KTransformers instance-scoped Conv3D fallback for torch 2.9.x VLM training.")
|
||||||
|
|
||||||
if not is_trainable:
|
if not is_trainable:
|
||||||
model.requires_grad_(False)
|
model.requires_grad_(False)
|
||||||
|
|||||||
@@ -51,8 +51,11 @@ def _get_gradient_checkpointing_kwargs(model_args: "ModelArguments") -> dict[str
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
from kt_kernel.sft import get_activation_checkpoint_context_fn
|
from kt_kernel.sft import get_activation_checkpoint_context_fn
|
||||||
except (ImportError, ModuleNotFoundError) as exc:
|
except (ImportError, ModuleNotFoundError):
|
||||||
raise RuntimeError("The installed kt-kernel does not provide the activation checkpoint context API.") from exc
|
logger.warning_rank0_once(
|
||||||
|
"The installed kt-kernel predates the activation checkpoint context API; using non-reentrant checkpointing."
|
||||||
|
)
|
||||||
|
return {"use_reentrant": False}
|
||||||
|
|
||||||
return {"use_reentrant": False, "context_fn": get_activation_checkpoint_context_fn()}
|
return {"use_reentrant": False, "context_fn": get_activation_checkpoint_context_fn()}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user