基于分层数据类的深度学习配置框架。
项目描述
微配置文件
用于深度学习的自以为是的 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。
运行演示:
- 导航到根目录
pip install -r requirements.txtexport PYTHONPATH="$PWD"cd scriptspython 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,AdamWConfiggeneral_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将仅在上述层次结构中加载一次,即使两者都LMModelConfig将TrainLoop其作为输入。
提供了一种解析命令行参数的方法。
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 的哈希值
| 算法 | 哈希摘要 | |
|---|---|---|
| SHA256 | 2074429cbc32d741d68c4712386305dbcb3f71b823a40e76ae1187f77cfc5ed7 |
|
| MD5 | 6bff3e6fe7f1c7294116aecd3a4d5715 |
|
| 布莱克2-256 | 2a098ab0973ef7d2ade2c9dfba3769eeef2fce0dc539cd501e5d55611bbcab89 |