Skip to main content

基于分层数据类的深度学习配置框架。

项目描述

微配置文件

用于深度学习的自以为是的 Python 数据类配置框架

我希望这种配置方法能让你的生活更轻松,但如果没有,请提交拉取请求或问题,我会看看我能做什么。这个配置系统肯定正在开发中,所以欢迎任何新的想法或建议。

安装

点安装:

pip install micro-config

或从源安装:

放置micro_config.py在项目的根目录。

回购指南

配置框架定义在micro_config.py.

repo 的其余部分提供了一个演示,说明人们可能希望如何micro_config.py在 pytorch 深度学习项目中实际使用。具体来说,我在 pytorch 的 wikitext 上实现了 Transformer 语言模型训练。有关如何micro_config.py在 jax(flax 或 haiku)中使用的另一个示例,请参阅此jax_v_pytorch repo

运行演示:

  1. 导航到根目录
  2. pip install -r requirements.txt
  3. export PYTHONPATH="$PWD"
  4. cd scripts
  5. python train_lm.py

您可以选择将命令行参数定义为train_lm.py

python train_lm.py epochs=1 bsize=16 model.transformer_config.hidden_dim=256

演示项目代码概览:

  • scripts/train_lm.py定义训练配置和脚本执行。
  • base_config.py为所有主要配置对象定义配置模式和默认值:WikiDataConfig, TransformerConfig, LMModelConfig,AdamWConfig
  • general_train_loop.py定义训练模型的配置模式和脚本。
  • src/定义了所有核心演示项目代码。

快速入门/演练

本节中的大部分演示代码均来自 repo 中提供的演示项目。

Python 数据类提供了比.yaml文件更自然、更灵活的配置定义接口。

  • 所有配置模式都应该被定义为一个实例ConfigScript或者ConfigScriptModel包含一个@dataclass装饰器
  • ConfigScripts 首先定义一个参数模式和可选的默认配置值。

例如,一个简单的数据集对象配置:

from dataclasses import dataclass, adsict
from micro_config import ConfigScript

# data config
@dataclass
class WikiDataConfig(ConfigScript):
    f_path: str='data/wikitext-2-raw/wiki.train.raw'
    max_len: int=256

ConfigScripts 加载相关的对象或函数。

  • 为此,所有人都ConfigScript实施unroll(self, metaconfig: MetaConfig).
  • metaconfig参数是另一个数据类,它为配置框架指定配置。随意子类化MetaConfig

例如,从配置加载数据集:

from dataclasses import dataclass, adsict
from micro_config import ConfigScript, MetaConfig
from src.data import WikitextDataset
import torch
import os

# data config
@dataclass
class WikiDataConfig(ConfigScript):
    f_path: str='data/wikitext-2-raw/wiki.train.raw'
    max_len: int=256

    def unroll(self, metaconfig: MetaConfig):
        # metaconfig.convert_path converts paths reletive to metaconfig.project_root into absolute paths
        return WikitextDataset(metaconfig.convert_path(self.f_path), self.max_len)

if __name__ == "__main__":
    metaconfig = MetaConfig(project_root=os.path.dirname(__file__), 
                            verbose=True)
    
    data_config = WikiDataConfig(max_len=512)
    data = data_config.unroll(metaconfig)

可以分层定义配置。

  • 您可以定义ConfigScripts为其他参数ConfigScripts
  • 您可以将ConfigScripts 的列表或字典定义为 a 的参数,方法是ConfigScript将您的列表或字典分别包装在ConfigScriptListorConfigScriptDict中。

例如,下面的 LM 模型配置将ConfigScripts 定义为数据集和 atransformer_config作为参数:

from micro_config import MetaConfig
from base_configs import ConfigScriptModel
from dataclasses import field
from src.lm import LMModel
import os

# model config
@dataclass
class LMModelConfig(ConfigScriptModel):
    dataset: WikiDataConfig=field(default_factory=lambda: WikiDataConfig())
    transformer_config: TransformerConfig=field(default_factory=lambda: TransformerConfig(max_len=256))

    def unroll(self, metaconfig: MetaConfig):
        dataset = self.dataset.unroll(metaconfig)
        transformer_config = self.transformer_config.unroll(metaconfig)
        return LMModel(dataset, transformer_config, self.device)

if __name__ == "__main__":
    metaconfig = MetaConfig(project_root=os.path.dirname(__file__), 
                            verbose=True)

    model_config = LMModelConfig(
        checkpoint_path=None, 
        strict_load=True, 
        device='cpu', 
        dataset=WikiDataConfig(f_path='data/wikitext-2-raw/wiki.train.raw', max_len=256), 
        transformer_config=TransformerConfig(
            max_length=256, 
            heads=12, 
            hidden_dim=768, 
            attn_dim=64, 
            intermediate_dim=3072, 
            num_blocks=12, 
            dropout=0.1
        )
    )
    model = model_config.unroll(metaconfig)

ConfigScriptModel(未提供micro_config开箱即用),如上所用,它是一个子类,ConfigScript它定义了一些默认功能,用于加载 unroll 返回的 pytorch 模块并将其放置在指定的设备上。您可以查看内部base_configs.py(或jax_v_pytorch 存储库)以查看如何实现此类特殊功能的示例。

配置和脚本是统一的:配置对脚本就像脚本对配置一样。

  • unroll(self, metaconfig: MetaConfig)不仅可以用来加载对象,还可以用来定义脚本逻辑。

例如,让我们定义一个简单的可配置训练循环:

from src.utils import combine_logs
from micro_config import ConfigScript, MetaConfig
from base_configs import ConfigScriptModel

@dataclass
class TrainLoop(ConfigScript):
    train_dataset: ConfigScript
    eval_dataset: ConfigScript
    model: ConfigScriptModel
    optim: ConfigScript
    epochs: int=10
    bsize: int=32
    
    def unroll(self, metaconfig: MetaConfig):
        print('using config:', asdict(self))
        device = metaconfig.device
        train_dataset = self.train_dataset.unroll(metaconfig)
        eval_dataset = self.eval_dataset.unroll(metaconfig)
        model = self.model.unroll(metaconfig)
        model.train()
        train_dataloader = DataLoader(train_dataset, batch_size=self.bsize)
        eval_dataloader = DataLoader(eval_dataset, batch_size=self.bsize)
        optim = self.optim.unroll(metaconfig)(model)
        for epoch in range(epochs):
            for x in tqdm(train_dataloader):
                loss, logs = model.get_loss(x.to(device))
                optim.zero_grad()
                loss.backward()
                optim.step()
                model.eval()
                val_x = next(iter(eval_dataloader))
                _, val_logs = model.get_loss(val_x.to(device))
                out_log = print({'train': combine_logs([logs]), 'val': combine_logs([val_logs]), 'step': (step+1)})
                model.train()
        return model

unroll(self, metaconfig: MetaConfig)根据配置层次结构的引用结构返回的对象。

  • 如果同一个配置对象在配置层次结构中被多次引用,则该对象的unroll(self, metaconfig: MetaConfig)方法将只被调用一次并且其输出被缓存,后续调用将返回缓存的输出。如果您不想要这种缓存行为,您可以ConfigScriptNoCache改为子类化。

例如,train_dataset在 中被引用两次train_config_script

import torch
import os

train_dataset = WikiDataConfig(f_path='data/wikitext-2-raw/wiki.train.raw', max_len=256)
eval_dataset = WikiDataConfig(f_path='data/wikitext-2-raw/wiki.valid.raw', max_len=256)

model = LMModelConfig(
            checkpoint_path=None, 
            strict_load=True, 
            device='cpu', 
            dataset=train_dataset, 
            transformer_config=TransformerConfig(
                max_length=256, 
                heads=12, 
                hidden_dim=768, 
                attn_dim=64, 
                intermediate_dim=3072, 
                num_blocks=12, 
                dropout=0.1
            )
        )

train_config_script = TrainLoop(
    train_dataset=train_dataset, 
    eval_dataset=eval_dataset, 
    model=model, 
    optim=AdamWConfig(lr=1e-4, weight_decay=0.01), 
    epochs=10, 
    bsize=16, 
)

if __name__ == "__main__":
    metaconfig = MetaConfig(project_root=os.path.dirname(__file__), 
                            verbose=True)
    # run the script
    train_config_script.unroll(metaconfig)

由 配置的数据集对象train_dataset将仅在上述层次结构中加载一次,即使两者都LMModelConfigTrainLoop其作为输入。

提供了一种解析命令行参数的方法。

  • parse_args(config)将命令行参数解析为字典
  • deep_replace(config, **kwargs)实现标准dataclasses.replace函数的嵌套版本
from micro_config import parse_args, deep_replace, MetaConfig
import os

if __name__ == "__main__":
    metaconfig = MetaConfig(project_root=os.path.dirname(__file__), 
                            verbose=True)
    train_config_script = deep_replace(train_config_script, **parse_args())
    # run the script
    train_config_script.unroll(metaconfig)

要通过命令行编辑层次结构中的任何参数,请像这样调用脚本:

python train_lm.py epochs=1 bsize=16 model.transformer_config.hidden_dim=256

项目详情


下载文件

下载适用于您平台的文件。如果您不确定要选择哪个,请了解有关安装包的更多信息。

源分布

micro_config-0.1.3.tar.gz (6.4 kB 查看哈希

已上传 source