已开启
`_target_` 与 `build()` 使用说明 #315
TonghanZhang创建于  8月5日
TonghanZhang
TonghanZhang成员
8月5日 创建

Hyper_model中支持使用_target_指定需要创建的callable对象,例如:

model:
  _target_: hyper_models._transformers.HyperAutoModelForCausalLM.from_pretrained
  pretrained_model_name_or_path: ./outputs/training_demo/tiny_llama
  torch_dtype: auto
  attn_implementation: sdpa
  force_hf: true
  
 dataloader:
  _target_: hyper_models.components.data.dataloader.DataLoader
  shuffle: true
  drop_last: true
  use_background_prefetcher: false

在解析yaml的时候,框架只会保存_target_指向的callable参数,不会进行实例化,例如model:
在解析的时候,框架只会按照_target_是否callable,_target_对应的函数/类签名进行校验并打包,

def from_pretrained(
    pretrained_model_name_or_path: str,
    ...
    torch_dtype="auto",
    attn_implementation="sdpa",
    force_hf: bool = False,
    **kwargs,
)

解析后打包:

Target(
    _target_=HyperAutoModelForCausalLM.from_pretrained,
    # yaml中的参数和def from_pretrain的参数对齐即可,定义的函数如果需要添加my_para,(写好函数之后)yaml直接在model下面添加对应的参数字段就可以了
    target_path=(
        "hyper_models._transformers."
        "HyperAutoModelForCausalLM.from_pretrained"
    ),
    pretrained_model_name_or_path="./outputs/training_demo/tiny_llama",
    torch_dtype="auto",
    attn_implementation="sdpa",
    force_hf=True,

    # from_pretrained() 的默认参数也会被保存
    distributed_setup=None,
    backend=None,
    peft_config=None,
    validate_placement=False,
    qat_config=None,
    fp8_config=None,
    compile_config=None,
    freeze_config=None,
)

然后在trainer中会执行:

self.model = self.config.model.build(
	distributed_setup=self.distributed_setup,
    peft_config=self.peft_config,
)

这个等价于执行:

transformers.AutoModelForCausalLM.from_pretrained(
    pretrained_model_name_or_path="./outputs/training_demo/tiny_llama",
    torch_dtype="auto",
    ...
    distributed_setup=self.distributed_setup,
    peft_config=self.peft_config,
)

这里的.build是Target类的实现方法,所以不需要在from_pretrained函数里面再尝试添加build的实现,
distributed_setup=self.distributed_setup, peft_config=self.peft_config,
是BaseTrainer里面运行时的参数,在这里传入。
如果需要额外传入新的参数,比如说‘self.dp_cp_mesh’(也是已经在trainer中定义好的),可以直接放进来(只要你的function from_pretrain能够接收)。

同理,以class AdamW为例:

optimizer:
  _target_: hyper_models.components.optim.optimizer.optimizer.AdamW
  lr: 1.0e-4
  weight_decay: 0.01
  betas: [0.9, 0.999]
  eps: 1.0e-8
  foreach: false

解析后也是导入并保存 AdamW class,按照 AdamW.init() 的签名校验 YAML 参数,此时不会创建 optimizer:
对应这个签名:

class AdamW:
    def __init__(
        self,
        *,
        model: nn.Module,
        lr: float = 1e-4,
        weight_decay: float = 0.01,
        betas: tuple[float, float] = (0.9, 0.999),
        eps: float = 1e-8,
        foreach: Optional[bool] = None,
    ):

解析后得到

config.optimizer = Target(
    _target_=AdamW,
    target_path=(
        "hyper_models.components.optim.optimizer.optimizer.AdamW"
    ),
    lr=0.0001,
    weight_decay=0.01,
    betas=(0.9, 0.999),
    eps=1e-8,
    foreach=False,
)

model虽然是必须参数,但是解析阶段允许暂时缺少,因为需要等待模型构造完成后由Trainer提供

optimizer = config.optimizer.build(model=self.model)

此时调用AdamW.init,之后创建整个类,当然,也不需要在class adamW下面def一个build函数。

如果需要接收其他的运行时参数,比如这个时候需要self.dataloader,以类似的方式穿进去就可以了,只要class的初始化支持传入。

__init__里面用到的所有参数仍然可以自己定义,例如这样:

class AdamW:
    def __init__(
        self,
        *,
        model: nn.Module,
        config: Any
    ):

然后传入一个字典,当然这就丧失了框架解析的用处。

likedislike
TonghanZhangTonghanZhang成员
8月5日 修改了issue 的描述
TonghanZhangTonghanZhang成员
8月5日 修改了issue 的描述
townwish4gitcode
8月6日 评论:

以配置文件中的optimizer部分为例:

optimizer:
  _target_: hyper_models.components.optim.optimizer.optimizer.AdamW
  lr: 1.0e-4
  weight_decay: 0.01

其会被自动解析为如下伪代码表示的类实例:

from functools import partial
from typing import Any, Callable

class ConfigWithBuildfn:
    def __init__(
        self,
        _target_: Callable[..., Any],
        **kwargs,
    ):
        self._target_ = _target_
        self._config_kwargs = kwargs

    def build(self, **runtime_kwargs: Any) -> Any:
        build_fn = partial(self._target_, **self._config_kwargs)
        return build_fn(**runtime_kwargs)

    def other_config_method(self):
        # Example of another method that can be added to the class
        # use self._config_kwargs as a configuration dictionary
        ...

则对于该解析得到的optimizer_config,其:

optimizer_config._config_kwargs = {
    "lr": 1.0e-4,
    "weight_decay": 0.01,
}
optimizer_config.build = partial(AdamW, **optimizer_config._config_kwargs)

后续调用optimizer_config.build(model)时,会实际执行

partial(AdamW, **optimizer_config._config_kwargs)(model)

即AdamW(model, **optimizer_config._config_kwargs)

,返回一个优化器实例。

likedislike