| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| """The definition of model engine. |
| |
| How to use: |
| model_engine = ModelEngine(model_args, is_train=True) |
| model_engine.processor: Get the tokenizer or multi-modal processor. |
| model_engine.renderer: Get the renderer. |
| model_engine.model_config: Get the model configuration. |
| model_engine.model: Get the HF model. |
| |
| Init workflow: |
| 1. Init processor. |
| 2. Init render. |
| 2. Init model config. |
| 3. Init model. |
| 4. Init adapter. |
| """ |
|
|
| import torch |
| from accelerate import init_empty_weights |
| from transformers import AutoConfig, AutoProcessor |
|
|
| from ..accelerator.helper import DeviceType |
| from ..accelerator.interface import DistributedInterface |
| from ..config.model_args import ModelArguments, ModelClass |
| from ..utils import logging |
| from ..utils.types import HFConfig, HFModel, Processor |
| from .utils.rendering import Renderer |
|
|
|
|
| logger = logging.get_logger(__name__) |
|
|
|
|
| class ModelEngine: |
| """Model engine. |
| |
| Args: |
| model_args: Model arguments. |
| is_train: Whether to train the model. |
| """ |
|
|
| def __init__( |
| self, |
| model_args: ModelArguments, |
| is_train: bool = False, |
| ) -> None: |
| self.args = model_args |
| """Model arguments.""" |
| self.is_train = is_train |
| """Whether to train the model.""" |
| self.processor = self._init_processor() |
| """Tokenizer or multi-modal processor.""" |
| self.renderer = Renderer(self.args.template, self.processor) |
| """Renderer.""" |
| self.model_config = self._init_model_config() |
| """Model configuration.""" |
| self._dist_config = DistributedInterface().dist_config |
| self._deepspeed_zero3_plugin = None |
| self._deepspeed_zero3_enabled = False |
|
|
| if self.is_train and self._dist_config is not None and self._dist_config.get("name") == "deepspeed": |
| from ..plugins.model_plugins.deepspeed_utils import ( |
| setup_deepspeed_zero3_model_loading, |
| teardown_deepspeed_zero3_model_loading, |
| ) |
|
|
| try: |
| self._deepspeed_zero3_plugin = setup_deepspeed_zero3_model_loading(self.is_train, self._dist_config) |
| self._deepspeed_zero3_enabled = self._deepspeed_zero3_plugin is not None |
| self.model = self._init_model() |
| finally: |
| teardown_deepspeed_zero3_model_loading(self._deepspeed_zero3_plugin) |
| self._deepspeed_zero3_plugin = None |
| self._deepspeed_zero3_enabled = False |
| else: |
| self.model = self._init_model() |
|
|
| def _init_processor(self) -> Processor: |
| """Init processor. |
| |
| NOTE: Transformers v5 always use fast tokenizer. |
| https://github.com/huggingface/transformers/blob/v5.0.0rc1/src/transformers/models/auto/tokenization_auto.py#L642 |
| """ |
| return AutoProcessor.from_pretrained( |
| self.args.model, |
| trust_remote_code=self.args.trust_remote_code, |
| ) |
|
|
| def _init_model_config(self) -> HFConfig: |
| """Init model config.""" |
| return AutoConfig.from_pretrained( |
| self.args.model, |
| trust_remote_code=self.args.trust_remote_code, |
| ) |
|
|
| def _init_model(self) -> HFModel: |
| """Init model. |
| |
| Transformers can choose the proper model init context. |
| https://github.com/huggingface/transformers/blob/v5.0.0rc0/src/transformers/modeling_utils.py#L3538 |
| """ |
| if self.args.init_config is not None: |
| from ..plugins.model_plugins.initialization import InitPlugin |
|
|
| init_device = InitPlugin(self.args.init_config.name)() |
| else: |
| init_device = DistributedInterface().current_device |
|
|
| init_kwargs = {} if self._deepspeed_zero3_enabled else {"device_map": init_device} |
| logger.info_rank0(f"Using attention implementation: {self.args.flash_attn}.") |
|
|
| if self.args.quant_config is not None: |
| from ..plugins.model_plugins.quantization import QuantizationPlugin |
|
|
| init_kwargs = QuantizationPlugin(self.args.quant_config.name)( |
| init_kwargs=init_kwargs, |
| config=self.model_config, |
| tokenizer=self.processor, |
| model_args=self.args, |
| is_trainable=self.is_train, |
| ) |
|
|
| if self.args.model_class == ModelClass.LLM: |
| from transformers import AutoModelForCausalLM, AutoModelForImageTextToText |
|
|
| if type(self.model_config) in AutoModelForImageTextToText._model_mapping.keys(): |
| AutoClass = AutoModelForImageTextToText |
| else: |
| AutoClass = AutoModelForCausalLM |
|
|
| elif self.args.model_class == ModelClass.CLS: |
| from transformers import AutoModelForTokenClassification |
|
|
| self.model_config.num_labels = 1 |
| self.model_config.classifier_dropout = 0.0 |
| text_config = getattr(self.model_config, "text_config", None) |
| if text_config is not None: |
| text_config.num_labels = 1 |
| text_config.classifier_dropout = 0.0 |
| AutoClass = AutoModelForTokenClassification |
| else: |
| from transformers import AutoModel |
|
|
| AutoClass = AutoModel |
|
|
| if init_device.type == DeviceType.META: |
| assert self.args.quant_config is None, "Quantization is not supported with meta device." |
| with init_empty_weights(): |
| model = AutoClass.from_config(self.model_config) |
| else: |
| model = AutoClass.from_pretrained( |
| self.args.model, |
| config=self.model_config, |
| dtype="auto", |
| attn_implementation=self.args.flash_attn, |
| trust_remote_code=self.args.trust_remote_code, |
| **init_kwargs, |
| ) |
|
|
| init_mode = self.args.init_config.name if self.args.init_config is not None else "init_on_default" |
| model._init_mode = init_mode |
|
|
| if self.args.peft_config is None: |
| if self.is_train: |
| logger.info_rank0("Fine-tuning mode: full tuning") |
| model = model.to(torch.float32) |
| else: |
| logger.info_rank0("Inference the original model") |
| else: |
| if self.args.peft_config.name == "lora" and init_mode == "init_on_meta": |
| raise ValueError("Currently lora stage does not support loading model by meta.") |
|
|
| from ..plugins.model_plugins.peft import PeftPlugin |
|
|
| model = PeftPlugin(self.args.peft_config.name)(model, self.args.peft_config, self.is_train) |
|
|
| if self.args.kernel_config is not None: |
| from ..plugins.model_plugins.kernels.interface import KernelPlugin |
|
|
| kernel_config = self.args.kernel_config |
| kernel_kwargs: dict = {"model": model, "include_kernels": kernel_config.get("include_kernels")} |
| if kernel_config.name == "liger_kernel": |
| |
| kernel_kwargs["require_logits"] = self.is_train |
| model = KernelPlugin(kernel_config.name)(**kernel_kwargs) |
|
|
| return model |
|
|
|
|
| if __name__ == "__main__": |
| """ |
| python -m llamafactory.v1.core.model_engine --model llamafactory/tiny-random-qwen2.5 |
| """ |
| from ..config.arg_parser import get_args |
|
|
| model_args, *_ = get_args() |
| model_engine = ModelEngine(model_args=model_args) |
| print(model_engine.processor) |
| print(model_engine.model_config) |
| print(model_engine.model) |
|
|