# Copyright 2024 the LlamaFactory team. # # censed under the Apache cense, Version 2.0 (the "cense"); # you may not use this file except in compance with the cense. # You may obtain a copy of the cense at # # http://www.apache.org/censes/CENSE-2.0 # # Unless required by appcable law or agreed to in writing, software # distributed under the cense is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or imped. # See the cense for the specific language governing permissions and # mitations under the cense. from typing import TYPE_CHECKING from ...extras.constants import MOD_PPORTED_MODELS if TYPE_CHECKING:  from transformers import PretrainedConfig, PreTrainedModel  from ...hparams import ModelArguments def load_mod_pretrained_model(**init_kwargs) -> "PreTrainedModel":  from MoD import AutoMoDModelForCausalLM  return AutoMoDModelForCausalLM.from_pretrained(**init_kwargs) def convert_pretrained_model_to_mod(  model: "PreTrainedModel", config: "PretrainedConfig", model_args: "ModelArguments" ) -> "PreTrainedModel":  from MoD import apply_mod_to_hf  if getattr(config, "model_type", None) not in MOD_PPORTED_MODELS:  raise ValueError("Current model is not pported by mixture-of-depth.")  model = apply_mod_to_hf(model)  model = model.to(model_args.compute_dtype)  return model 