mirror of
https://github.com/PaddlePaddle/PaddleOCR.git
synced 2026-09-24 23:33:08 +08:00
add benckmark
This commit is contained in:
+27
-10
@@ -154,6 +154,24 @@ def check_xpu(use_xpu):
|
||||
except Exception as e:
|
||||
pass
|
||||
|
||||
def to_float32(preds):
|
||||
if isinstance(preds, dict):
|
||||
for k in preds:
|
||||
if isinstance(preds[k], dict) or isinstance(preds[k], list):
|
||||
preds[k] = to_float32(preds[k])
|
||||
else:
|
||||
preds[k] = preds[k].astype(paddle.float32)
|
||||
elif isinstance(preds, list):
|
||||
for k in range(len(preds)):
|
||||
if isinstance(preds[k], dict):
|
||||
preds[k] = to_float32(preds[k])
|
||||
elif isinstance(preds[k], list):
|
||||
preds[k] = to_float32(preds[k])
|
||||
else:
|
||||
preds[k] = preds[k].astype(paddle.float32)
|
||||
else:
|
||||
preds = preds.astype(paddle.float32)
|
||||
return preds
|
||||
|
||||
def train(config,
|
||||
train_dataloader,
|
||||
@@ -252,13 +270,19 @@ def train(config,
|
||||
|
||||
# use amp
|
||||
if scaler:
|
||||
with paddle.amp.auto_cast():
|
||||
with paddle.amp.auto_cast(level='O2'):
|
||||
if model_type == 'table' or extra_input:
|
||||
preds = model(images, data=batch[1:])
|
||||
elif model_type in ["kie", 'vqa']:
|
||||
preds = model(batch)
|
||||
else:
|
||||
preds = model(images)
|
||||
preds = to_float32(preds)
|
||||
loss = loss_class(preds, batch)
|
||||
avg_loss = loss['loss']
|
||||
scaled_avg_loss = scaler.scale(avg_loss)
|
||||
scaled_avg_loss.backward()
|
||||
scaler.minimize(optimizer, scaled_avg_loss)
|
||||
else:
|
||||
if model_type == 'table' or extra_input:
|
||||
preds = model(images, data=batch[1:])
|
||||
@@ -266,15 +290,8 @@ def train(config,
|
||||
preds = model(batch)
|
||||
else:
|
||||
preds = model(images)
|
||||
|
||||
loss = loss_class(preds, batch)
|
||||
avg_loss = loss['loss']
|
||||
|
||||
if scaler:
|
||||
scaled_avg_loss = scaler.scale(avg_loss)
|
||||
scaled_avg_loss.backward()
|
||||
scaler.minimize(optimizer, scaled_avg_loss)
|
||||
else:
|
||||
loss = loss_class(preds, batch)
|
||||
avg_loss = loss['loss']
|
||||
avg_loss.backward()
|
||||
optimizer.step()
|
||||
optimizer.clear_grad()
|
||||
|
||||
Reference in New Issue
Block a user