mirror of
https://github.com/PaddlePaddle/PaddleOCR.git
synced 2026-09-24 23:33:08 +08:00
add pse to benchmark
This commit is contained in:
+18
-12
@@ -212,15 +212,15 @@ def train(config,
|
||||
for epoch in range(start_epoch, epoch_num + 1):
|
||||
train_dataloader = build_dataloader(
|
||||
config, 'Train', device, logger, seed=epoch)
|
||||
train_batch_cost = 0.0
|
||||
train_reader_cost = 0.0
|
||||
batch_sum = 0
|
||||
batch_start = time.time()
|
||||
train_run_cost = 0.0
|
||||
total_samples = 0
|
||||
reader_start = time.time()
|
||||
max_iter = len(train_dataloader) - 1 if platform.system(
|
||||
) == "Windows" else len(train_dataloader)
|
||||
for idx, batch in enumerate(train_dataloader):
|
||||
profiler.add_profiler_step(profiler_options)
|
||||
train_reader_cost += time.time() - batch_start
|
||||
train_reader_cost += time.time() - reader_start
|
||||
if idx >= max_iter:
|
||||
break
|
||||
lr = optimizer.get_lr()
|
||||
@@ -228,6 +228,7 @@ def train(config,
|
||||
if use_srn:
|
||||
model_average = True
|
||||
|
||||
train_start = time.time()
|
||||
# use amp
|
||||
if scaler:
|
||||
with paddle.amp.auto_cast():
|
||||
@@ -252,8 +253,8 @@ def train(config,
|
||||
optimizer.step()
|
||||
optimizer.clear_grad()
|
||||
|
||||
train_batch_cost += time.time() - batch_start
|
||||
batch_sum += len(images)
|
||||
train_run_cost += time.time() - train_start
|
||||
total_samples += len(images)
|
||||
|
||||
if not isinstance(lr_scheduler, float):
|
||||
lr_scheduler.step()
|
||||
@@ -284,12 +285,13 @@ def train(config,
|
||||
logs = train_stats.log()
|
||||
strs = 'epoch: [{}/{}], iter: {}, {}, reader_cost: {:.5f} s, batch_cost: {:.5f} s, samples: {}, ips: {:.5f}'.format(
|
||||
epoch, epoch_num, global_step, logs, train_reader_cost /
|
||||
print_batch_step, train_batch_cost / print_batch_step,
|
||||
batch_sum, batch_sum / train_batch_cost)
|
||||
print_batch_step, (train_reader_cost + train_run_cost) /
|
||||
print_batch_step, total_samples,
|
||||
total_samples / (train_reader_cost + train_run_cost))
|
||||
logger.info(strs)
|
||||
train_batch_cost = 0.0
|
||||
train_reader_cost = 0.0
|
||||
batch_sum = 0
|
||||
train_run_cost = 0.0
|
||||
total_samples = 0
|
||||
# eval
|
||||
if global_step > start_eval_step and \
|
||||
(global_step - start_eval_step) % eval_batch_step == 0 and dist.get_rank() == 0:
|
||||
@@ -342,7 +344,7 @@ def train(config,
|
||||
global_step)
|
||||
global_step += 1
|
||||
optimizer.clear_grad()
|
||||
batch_start = time.time()
|
||||
reader_start = time.time()
|
||||
if dist.get_rank() == 0:
|
||||
save_model(
|
||||
model,
|
||||
@@ -383,7 +385,11 @@ def eval(model,
|
||||
with paddle.no_grad():
|
||||
total_frame = 0.0
|
||||
total_time = 0.0
|
||||
pbar = tqdm(total=len(valid_dataloader), desc='eval model:')
|
||||
pbar = tqdm(
|
||||
total=len(valid_dataloader),
|
||||
desc='eval model:',
|
||||
position=0,
|
||||
leave=True)
|
||||
max_iter = len(valid_dataloader) - 1 if platform.system(
|
||||
) == "Windows" else len(valid_dataloader)
|
||||
for idx, batch in enumerate(valid_dataloader):
|
||||
|
||||
Reference in New Issue
Block a user