未验证 提交 95cf3a64 编写于 作者: D dyning 提交者: GitHub

Merge pull request #89 from WuHaobo/master

polish
...@@ -48,8 +48,6 @@ TRAIN: ...@@ -48,8 +48,6 @@ TRAIN:
order: '' order: ''
- ToCHWImage: - ToCHWImage:
VALID: VALID:
batch_size: 64 batch_size: 64
num_workers: 4 num_workers: 4
...@@ -72,4 +70,3 @@ VALID: ...@@ -72,4 +70,3 @@ VALID:
order: '' order: ''
- ToCHWImage: - ToCHWImage:
...@@ -48,8 +48,6 @@ TRAIN: ...@@ -48,8 +48,6 @@ TRAIN:
order: '' order: ''
- ToCHWImage: - ToCHWImage:
VALID: VALID:
batch_size: 64 batch_size: 64
num_workers: 4 num_workers: 4
...@@ -72,4 +70,3 @@ VALID: ...@@ -72,4 +70,3 @@ VALID:
order: '' order: ''
- ToCHWImage: - ToCHWImage:
...@@ -48,4 +48,25 @@ TRAIN: ...@@ -48,4 +48,25 @@ TRAIN:
order: '' order: ''
- ToCHWImage: - ToCHWImage:
VALID:
batch_size: 64
num_workers: 4
file_list: "./dataset/ILSVRC2012/val_list.txt"
data_dir: "./dataset/ILSVRC2012/"
shuffle_seed: 0
transforms:
- DecodeImage:
to_rgb: True
to_np: False
channel_first: False
- ResizeImage:
resize_short: 256
- CropImage:
size: 224
- NormalizeImage:
scale: 1.0/255.0
mean: [0.485, 0.456, 0.406]
std: [0.229, 0.224, 0.225]
order: ''
- ToCHWImage:
...@@ -95,7 +95,7 @@ def main(args): ...@@ -95,7 +95,7 @@ def main(args):
# 1. train with train dataset # 1. train with train dataset
program.run(train_dataloader, exe, compiled_train_prog, train_fetchs, program.run(train_dataloader, exe, compiled_train_prog, train_fetchs,
epoch_id, 'train') epoch_id, 'train')
if int(os.environ.get("PADDLE_TRAINERS_ID", 0)) == 0: if int(os.getenv("PADDLE_TRAINER_ID", 0)) == 0:
# 2. validate with validate dataset # 2. validate with validate dataset
if config.validate and epoch_id % config.valid_interval == 0: if config.validate and epoch_id % config.valid_interval == 0:
top1_acc = program.run(valid_dataloader, exe, top1_acc = program.run(valid_dataloader, exe,
......
Markdown is supported
0% .
You are about to add 0 people to the discussion. Proceed with caution.
先完成此消息的编辑!
想要评论请 注册