使用遗传算法训练 keras 模型。
项目描述
凯拉斯基因
Keras Genetic 允许您使用遗传算法轻松训练 Keras 模型。
快速链接:
背景
KerasGenetic 允许您在使用遗传算法进行训练时利用优雅的建模 API Keras。通常,Keras 神经网络权重是通过梯度下降过程最小化损失函数来优化的。
Keras Genetic 通过利用遗传算法采用不同的方法进行权重优化。遗传算法允许您在没有关于损失情况信息的情况下优化神经网络。
遗传算法可用于在需要训练具有 <1000 个参数的专用控制器的特殊情况下训练神经网络。
今天应用遗传算法的一些领域:
概述
Keras 遗传 API 上手速度很快,但足够灵活,可以适应您可能提出的任何用例。
必须使用 API 的三个核心组件才能开始:
- 这
Individual - 这
Evaluator - 这
Breeder search()
个人
类代表群体中的Individual个体。
类中最重要的方法Individual是load_model().
load_model()生成一个 Keras 模型,其权重存储在individual
加载的类中:
model = individual.load_model()
model.predict(some_data)
评估者
接下来,让我们回顾一下Evaluator。Evaluator负责确定一个的强度Individual。也许最简单的评估器是分类任务的准确度评估器:
def evaluate_accuracy(individual: keras_genetic.Individual):
model = individual.load_model()
result = model.evaluate(x_train[:100], y_train[:100], return_dict=True, verbose=0)
return result["accuracy"]
上面定义的evaluate_accuracy()函数将 a 映射Individual到准确度分数。该分数可用于选择将在下一代中使用的个体。
饲养员
Breeder负责从一组父个体中产生新个体。每个人如何Breeder产生新个体的细节对于育种者来说是独一无二的,但作为一般规则,父母的一些属性被保留,而一些新属性是随机抽样的。
对于大多数用户来说,这TwoParentMutationBreeder已经足够有效了。
搜索()
search()类似于model.fit()核心 Keras 框架。search()API 支持多种参数。如需深入了解,请浏览 API 文档。
以下是该search()函数的示例用法:
results = keras_genetic.search(
model=model,
# computational cost is evaluate*generations*population_size
evaluator=evaluate_accuracy,
generations=10,
population_size=50,
n_parents_from_population=5,
breeder=keras_genetic.breeder.MutationBreeder(),
return_best=1,
)
延伸阅读
查看示例和指南(即将推出!)。
快速开始
目前,Cartpole 示例用作快速入门指南。
路线图
我想完成以下任务:
- ✅ 稳定基础 API
- ✅ 支持回调 API
- ✅ 端到端 MNIST 示例
- ✅ 端到端 CartPole 示例
- ✅ 实现一个 ProgBarLogger
- ✅ 实现 EarlyStopping 回调(可用于 CartPole 示例)
- 至少有 3 个不同的育种者
- 自动生成文档
- 彻底记录每个组件
- 提供最有效的遗传算法的实现
- 为每个组件实施单元测试
- 支持随机播种
- 根据 Keras 核心 API 设计指南彻底审查 API
- 支持自定义初始种群(即模仿人类模仿模型)
- 支持 keep_probability 计划
随意贡献任何这些。
引文
@misc{wood2022kerasgenetic,
title = {Keras Genetic},
author = {Luke Wood},
year = 2022,
howpublished = {\url{https://github.com/lukewood/keras-genetic}}
}
项目详情
下载文件
下载适用于您平台的文件。如果您不确定要选择哪个,请了解有关安装包的更多信息。
源分布
内置分布
keras_genetic -0.1.0-py3-none-any.whl 的哈希值
| 算法 | 哈希摘要 | |
|---|---|---|
| SHA256 | 53d259ee7ff3a58d94249a82001e176494aa06611a8f34b898b8d48d27446a4d |
|
| MD5 | e012d8005224c0ca8049f3f3d76ddbb5 |
|
| 布莱克2-256 | 958d948e9d7db69cea23d211f63f2f9b8c737ff662e29e695ac71da2a68d7f6d |