mirror of
https://github.com/PaddlePaddle/PaddleOCR.git
synced 2026-09-24 23:33:08 +08:00
Merge pull request #7139 from andyjpaddle/fix_amp_re
[TIPC] Fix amp train for re
This commit is contained in:
+5
-3
@@ -154,13 +154,14 @@ 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)
|
||||
preds[k] = paddle.to_tensor(preds[k], dtype='float32')
|
||||
elif isinstance(preds, list):
|
||||
for k in range(len(preds)):
|
||||
if isinstance(preds[k], dict):
|
||||
@@ -168,11 +169,12 @@ def to_float32(preds):
|
||||
elif isinstance(preds[k], list):
|
||||
preds[k] = to_float32(preds[k])
|
||||
else:
|
||||
preds[k] = preds[k].astype(paddle.float32)
|
||||
preds[k] = paddle.to_tensor(preds[k], dtype='float32')
|
||||
else:
|
||||
preds = preds.astype(paddle.float32)
|
||||
preds = paddle.to_tensor(preds, dtype='float32')
|
||||
return preds
|
||||
|
||||
|
||||
def train(config,
|
||||
train_dataloader,
|
||||
valid_dataloader,
|
||||
|
||||
Reference in New Issue
Block a user