File size: 11,680 Bytes
07cb7d3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
# 使用教程

这份教程帮助你快速跑通项目,并在现有任务的基础上替换数据、替换模型或新增任务。

## 快速入门:跑通图片分类任务

建议先从 `tasks/image_classification/runner.py` 开始。图片分类任务最直观,只要准备好猫狗图片目录,就可以完整体验训练、样例验证和模型导出。

### 1. 准备环境

本地开发使用 `environment.yml````bash
conda env create -f environment.yml
conda activate general-dl
```

如果环境已经存在,可以更新:

```bash
conda env update -f environment.yml --prune
conda activate general-dl
```

Linux 服务器或正式训练环境使用 `environment-linux.yml`### 2. 准备数据

图片分类任务默认读取这些目录:

```text
~/data/cat-vs-dog/PetImagesMini/train
~/data/cat-vs-dog/PetImagesMini/val
~/data/cat-vs-dog/PetImagesMini/test
```

目录应按 Keras 图片分类数据集格式组织,也就是每个类别一个子目录:

```text
train/
  Cat/
    cat1.jpg
  Dog/
    dog1.jpg
val/
  Cat/
  Dog/
test/
  Cat/
  Dog/
```

如果你的数据放在别处,修改 `tasks/image_classification/runner.py``CatsVsDogsDataSource``train_path``validation_path``test_path`### 3. 选择运行环境

项目用 `ENV` 环境变量选择配置:

```bash
ENV=dev
ENV=prod
```

`ENV=dev` 会选择开发配置,通常是小数据、小模型、少轮数,用来快速验证。`ENV=prod` 会选择正式配置,通常训练更久、模型更大。

如果不设置 `ENV`,默认使用正式配置。

### 4. 运行训练

```bash
ENV=dev python -m tasks.image_classification.runner train
```

训练时会生成:

```text
local/tasks/image_classification/logs/config.txt
local/tasks/image_classification/checkpoints/model_epoch_001.weights.h5
local/tasks/image_classification/tensorboard/
```

其中 `logs/config.txt` 会记录本次任务配置,`checkpoints` 保存训练检查点,`tensorboard` 保存 TensorBoard 日志。

### 5. 使用训练检查点跑样例

如果你刚训练完,还没有导出模型,可以用训练检查点跑样例:

```bash
ENV=dev python -m tasks.image_classification.runner test checkpoint=train
```

这个动作会加载 `local/tasks/image_classification/checkpoints/` 下的训练检查点,并打印若干样例的真实类别、预测类别和置信度。

### 6. 导出模型

```bash
ENV=dev python -m tasks.image_classification.runner export
```

导出的 Keras 模型会保存到:

```text
local/saved/models/image_classification/
```

正式环境的导出目录是:

```text
saved/models/image_classification/
```

### 7. 使用导出模型跑样例

```bash
ENV=dev python -m tasks.image_classification.runner test
```

不传 `checkpoint=train` 时,`test` 默认读取已导出的 `.keras` 模型。

## 流水线教程:任务脚本应该长什么样

`tasks/*/runner.py` 是每个任务的配置中心。它通常做四件事:

- 创建开发配置
- 创建正式配置
-`resolve_env(dev_conf, prod_conf)``ENV` 选择配置
- 把选中的 Pipeline 交给 `PipelineRunner`

图片分类任务的核心结构如下:

```python
pipeline = resolve_env(
    SupervisedModelPipeline(
        name="image_classification",
        data_source=...,
        model_builder=...,
        training_rule=...
    ),
    SupervisedModelPipeline(
        name="image_classification",
        data_source=...,
        model_builder=...,
        training_rule=...,
        checkpoint_load_rules=...
    )
)

pipeline_runner = PipelineRunner(pipeline)

if __name__ == "__main__":
    pipeline_runner()
```

`PipelineRunner` 暴露三个动作:

- `train`:训练模型
- `export`:从检查点导出 `.keras` 模型
- `test`:运行样例验证

`name` 很重要,它会影响产物路径。以 `name="image_classification"` 为例:

```text
local/tasks/image_classification/logs/
local/tasks/image_classification/checkpoints/
local/tasks/image_classification/tensorboard/
local/saved/models/image_classification/
```

有监督任务使用 `SupervisedModelPipeline`,例如图片分类、图像分割。文本生成任务使用 `TextGenerationPipeline`,例如诗歌生成和 Wiki GPT。

## 数据源教程:Pipeline 期待你提供什么

数据源负责把你的原始数据变成 Pipeline 能使用的数据。

### 有监督任务数据源

有监督任务的数据源需要提供:

- `training_ds()`:返回训练集和验证集
- `test_examples(model)`:使用模型跑几个样例,并打印或保存结果

最小形状如下:

```python
class MyDataSource:
    def training_ds(self):
        train_ds = ...
        validation_ds = ...
        return train_ds, validation_ds

    def test_examples(self, model):
        examples_ds = ...
        for images, labels in examples_ds:
            predictions = model.predict(images)
            print(predictions)
```

`training_ds()` 返回的两个对象会直接传给 Keras 的 `model.fit()``test_examples(model)` 接收的是可用于推理的 Keras 模型。

### 文本生成任务数据源

文本生成任务的数据源需要提供:

- `doc_ds()`:返回原始文本数据
- `tokens_ds()`:返回训练用 token 数据
- `tokenizer_bundle()`:返回分词器、反解码函数、词表大小和结束标记等推理资源

最小形状如下:

```python
class MyTextDataSource:
    data_dir = "data/my_text"
    sequence_length = 100
    batch_size = 32
    validation_batches = 0

    def doc_ds(self):
        return docs_ds

    def tokens_ds(self):
        return tokens_ds

    def tokenizer_bundle(self):
        return tokenizer_bundle
```

`tokens_ds()` 返回的数据应能被文本模型训练使用,通常每个元素是 `(input_ids, target_ids)`。批次大小和验证集切分相关配置统一放在数据源中配置。

## 模型构建教程:Pipeline 期待模型构建器做什么

模型构建器负责创建训练模型、把训练产物转换成推理产物,以及从完整模型文件加载推理产物。

### 有监督模型构建器

有监督模型构建器需要提供:

- `build_training_artifact()`
- `compile_training_model(model)`
- `convert_to_inference_artifact(training_artifact)`
- `load_inference_artifact(model_path)`

最小形状如下:

```python
class MyModelBuilder:
    def build_training_artifact(self):
        model = build_model()
        return ModelArtifact(model=model)

    def compile_training_model(self, model):
        model.compile(
            optimizer="adam",
            loss="binary_crossentropy",
            metrics=["accuracy"]
        )

    def convert_to_inference_artifact(self, training_artifact):
        return training_artifact

    def load_inference_artifact(self, model_path):
        model = keras.models.load_model(str(model_path))
        return ModelArtifact(model=model)
```

训练模型和推理模型可以相同,也可以不同。新手接入时可以先让 `convert_to_inference_artifact()` 直接返回 `training_artifact`### 文本生成模型构建器

文本生成模型构建器需要提供:

- `build_training_artifact(vocab_size, sequence_length)`
- `compile_training_model(model)`
- `convert_to_inference_artifact(training_artifact)`
- `load_inference_artifact(model_path)`

最小形状如下:

```python
class MyTextModelBuilder:
    def build_training_artifact(self, vocab_size, sequence_length):
        model = build_text_model(vocab_size, sequence_length)
        return TextGenerationModel(model=model, generate=generate_fn)

    def compile_training_model(self, model):
        model.compile(...)

    def convert_to_inference_artifact(self, training_artifact):
        return training_artifact

    def load_inference_artifact(self, model_path):
        model = keras.models.load_model(str(model_path))
        return TextGenerationModel(model=model, generate=generate_fn)
```

文本任务训练时会从数据源的 `tokenizer_bundle()` 里读取 `vocab_size`,从数据源读取 `sequence_length`,再交给模型构建器。

## 从零接入一个新任务

推荐先复制一个已有任务,再逐步替换。

### 1. 新建数据源

例如:

```text
src/deep_learning/data/my_task/dataset.py
```

如果是有监督任务,实现 `training_ds()``test_examples(model)`。如果是文本生成任务,实现 `doc_ds()``tokens_ds()``tokenizer_bundle()`### 2. 新建模型构建器

例如:

```text
src/deep_learning/models/my_task.py
```

简单模型可以直接放在 `src/deep_learning/models/my_task.py`,模型较复杂时再建目录拆分。无论是哪种任务,都把模型构建、训练模型编译放在模型构建器里。

### 3. 新建任务入口

例如:

```text
tasks/my_task/runner.py
```

先准备开发配置:

```python
SupervisedModelPipeline(
    name="my_task",
    data_source=MyDataSource(...),
    model_builder=MyModelBuilder(...),
    training_rule=TrainingRule(
        epochs=1,
        steps_per_epoch=1
    )
)
```

再准备正式配置,把数据路径、模型规模和训练轮数替换成正式训练需要的值。

### 4. 先跑开发训练

```bash
ENV=dev python -m tasks.my_task.runner train
```

### 5. 用训练检查点跑样例

```bash
ENV=dev python -m tasks.my_task.runner test checkpoint=train
```

### 6. 导出模型

```bash
ENV=dev python -m tasks.my_task.runner export
```

### 7. 用导出模型跑样例

```bash
ENV=dev python -m tasks.my_task.runner test
```

## 常见问题

### 检查点保存在哪里?

开发环境保存在:

```text
local/tasks/<task_name>/checkpoints/
```

正式环境也会使用任务运行目录:

```text
local/tasks/<task_name>/checkpoints/
```

### 导出模型保存在哪里?

开发环境保存在:

```text
local/saved/models/<task_name>/
```

正式环境保存在:

```text
saved/models/<task_name>/
```

### 日志保存在哪里?

任务配置日志保存在:

```text
local/tasks/<task_name>/logs/config.txt
```

TensorBoard 日志保存在:

```text
local/tasks/<task_name>/tensorboard/
```

启动 TensorBoard:

```bash
tensorboard --logdir=local/tasks/<task_name>/tensorboard
```

### 为什么同时需要开发配置和正式配置?

开发配置用于快速确认数据、模型和流程都能跑通。正式配置用于完整训练。这样可以先用很小的成本发现路径、形状和配置问题。

### 如何指定加载第几轮 checkpoint?

在任务入口里配置 `CheckpointLoadRules`。例如图片分类正式配置中:

```python
checkpoint_load_rules=CheckpointLoadRules(
    export=CheckpointConfig(epoch=13),
    test=CheckpointConfig(dirs=[resolve_saved("models/image_classification")], suffix=".keras")
)
```

这表示导出时优先加载第 13 轮训练检查点,测试时从已导出的模型目录读取 `.keras` 文件。

### `batch_size` 应该放在哪里?

`batch_size` 和验证集切分相关配置统一放在数据源中配置。数据源决定数据如何生成,`TrainingRule` 只描述训练跑多久,例如 `epochs` 和 `steps_per_epoch`。

### 没有 checkpoint 时哪些动作会失败?

`train` 可以在没有检查点时从新模型开始训练。`export` 需要训练检查点。`test` 默认读取导出模型,如果没有导出模型会失败;如果传 `checkpoint=train`,它会读取训练检查点。

### 文本任务为什么需要词表?

文本模型训练的是 token id,不是原始字符串。词表用来把文本转成 token id,也用来把生成结果从 token id 转回文字。文本任务的数据源需要通过 `tokenizer_bundle()` 提供这些资源。