Spaces:
Sleeping
Sleeping
| from deep_learning.data.cats_vs_dogs import CatsVsDogsDataSource | |
| from deep_learning.env.resolve import resolve_env, resolve_path, resolve_saved | |
| from deep_learning.models.image_classification import ImageClassificationModelBuilder | |
| from deep_learning.pipeline import ( | |
| SupervisedModelPipeline, | |
| PipelineRunner | |
| ) | |
| from deep_learning.pipeline.specs.configs import CheckpointConfig, CheckpointLoadRules, TrainingRule | |
| pipeline = resolve_env( | |
| # 开发配置 | |
| SupervisedModelPipeline( | |
| name="image_classification", | |
| data_source=CatsVsDogsDataSource( | |
| train_path=resolve_path("~/data/cat-vs-dog/PetImagesMini/train"), | |
| validation_path=resolve_path("~/data/cat-vs-dog/PetImagesMini/val"), | |
| test_path=resolve_path("~/data/cat-vs-dog/PetImagesMini/test"), | |
| image_size=(180, 180), | |
| label_mode="binary", | |
| batch_size=2, | |
| example_count=5 | |
| ), | |
| model_builder=ImageClassificationModelBuilder( | |
| image_size=(180, 180), | |
| model_filters=(32,) | |
| ), | |
| training_rule=TrainingRule( | |
| epochs=1, | |
| steps_per_epoch=1 | |
| ) | |
| ), | |
| # 生产配置 | |
| SupervisedModelPipeline( | |
| name="image_classification", | |
| data_source=CatsVsDogsDataSource( | |
| train_path=resolve_path("~/data/cat-vs-dog/PetImagesMini/train"), | |
| validation_path=resolve_path("~/data/cat-vs-dog/PetImagesMini/val"), | |
| test_path=resolve_path("~/data/cat-vs-dog/PetImagesMini/test"), | |
| image_size=(180, 180), | |
| label_mode="binary", | |
| batch_size=32, | |
| example_count=5 | |
| ), | |
| model_builder=ImageClassificationModelBuilder( | |
| image_size=(180, 180), | |
| model_filters=(128, 256, 512, 728) | |
| ), | |
| training_rule=TrainingRule( | |
| epochs=30, | |
| steps_per_epoch=None | |
| ), | |
| checkpoint_load_rules=CheckpointLoadRules( | |
| export=CheckpointConfig(epoch=13), | |
| test=CheckpointConfig(dirs=[resolve_saved("models/image_classification")], suffix=".keras") | |
| ) | |
| ) | |
| ) | |
| pipeline_runner = PipelineRunner(pipeline) | |
| if __name__ == "__main__": | |
| pipeline_runner() | |