Spaces:
Running on Zero
Running on Zero
Upload 5 files
Browse files- README.md +14 -0
- app.py +1 -1
- requirements.txt +1 -1
- training/train.py +37 -1
README.md
CHANGED
|
@@ -123,6 +123,20 @@ For an initial test on a small dataset, use fewer epochs such as 2–5. Once eve
|
|
| 123 |
|
| 124 |
A GPU Space is strongly recommended.
|
| 125 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 126 |
## Counting
|
| 127 |
|
| 128 |
After training, open the **Count** tab and upload one image.
|
|
|
|
| 123 |
|
| 124 |
A GPU Space is strongly recommended.
|
| 125 |
|
| 126 |
+
|
| 127 |
+
## Training fix
|
| 128 |
+
|
| 129 |
+
This release includes a CUDA-device fix for RT-DETR's contrastive-denoising training path. On some
|
| 130 |
+
Transformers/PyTorch combinations, the denoising class-index tensor can remain on CPU while the
|
| 131 |
+
RT-DETR class embedding is on CUDA, producing:
|
| 132 |
+
|
| 133 |
+
`RuntimeError: Expected all tensors to be on the same device ... cpu ... cuda:0`
|
| 134 |
+
|
| 135 |
+
The training script now moves nested target tensors explicitly and patches the RT-DETR denoising
|
| 136 |
+
helper so its target tensors follow the class-embedding device. The `num_labels=10` vs. checkpoint
|
| 137 |
+
`80` message is expected when fine-tuning the COCO-pretrained checkpoint for 10 custom classes;
|
| 138 |
+
`ignore_mismatched_sizes=True` intentionally reinitializes the classification heads.
|
| 139 |
+
|
| 140 |
## Counting
|
| 141 |
|
| 142 |
After training, open the **Count** tab and upload one image.
|
app.py
CHANGED
|
@@ -26,7 +26,7 @@ else:
|
|
| 26 |
ROOT = Path(__file__).resolve().parent
|
| 27 |
|
| 28 |
# Self-heal the training script if a deployment omitted the training/ directory.
|
| 29 |
-
_TRAINING_SCRIPT_B64 = "aW1wb3J0IGFyZ3BhcnNlCmltcG9ydCBqc29uCmZyb20gcGF0aGxpYiBpbXBvcnQgUGF0aAoKaW1wb3J0IHRvcmNoCmZyb20gUElMIGltcG9ydCBJbWFnZQpmcm9tIHRvcmNoLnV0aWxzLmRhdGEgaW1wb3J0IERhdGFzZXQsIERhdGFMb2FkZXIKZnJvbSB0cWRtIGltcG9ydCB0cWRtCmZyb20gdHJhbnNmb3JtZXJzIGltcG9ydCBSVERldHJJbWFnZVByb2Nlc3NvciwgUlREZXRyRm9yT2JqZWN0RGV0ZWN0aW9uCgpCQVNFX01PREVMID0gIlBla2luZ1UvcnRkZXRyX3I1MHZkIgoKZGVmIGxvYWRfY2xhc3NlcyhwYXRoKToKICAgIHJldHVybiBbeC5zdHJpcCgpIGZvciB4IGluIFBhdGgocGF0aCkucmVhZF90ZXh0KCkuc3BsaXRsaW5lcygpIGlmIHguc3RyaXAoKV0KCmNsYXNzIENPQ09EZXRlY3Rpb25EYXRhc2V0KERhdGFzZXQpOgogICAgZGVmIF9faW5pdF9fKHNlbGYsIGltYWdlX2RpciwgYW5ub3RhdGlvbl9maWxlLCBwcm9jZXNzb3IpOgogICAgICAgIHNlbGYuaW1hZ2VfZGlyID0gUGF0aChpbWFnZV9kaXIpCiAgICAgICAgc2VsZi5wcm9jZXNzb3IgPSBwcm9jZXNzb3IKICAgICAgICBjb2NvID0ganNvbi5sb2FkcyhQYXRoKGFubm90YXRpb25fZmlsZSkucmVhZF90ZXh0KCkpCiAgICAgICAgc2VsZi5pbWFnZXMgPSB7eFsiaWQiXTogeCBmb3IgeCBpbiBjb2NvWyJpbWFnZXMiXX0KICAgICAgICBjYXRzID0gc29ydGVkKGNvY29bImNhdGVnb3JpZXMiXSwga2V5PWxhbWJkYSB4OiB4WyJpZCJdKQogICAgICAgIHNlbGYuY2F0ZWdvcnlfaWRfdG9fbGFiZWwgPSB7Y1siaWQiXTogaSBmb3IgaSxjIGluIGVudW1lcmF0ZShjYXRzKX0KICAgICAgICBhbm5zID0ge30KICAgICAgICBmb3IgYSBpbiBjb2NvWyJhbm5vdGF0aW9ucyJdOgogICAgICAgICAgICBpZiBub3QgYS5nZXQoImlzY3Jvd2QiLCAwKToKICAgICAgICAgICAgICAgIGFubnMuc2V0ZGVmYXVsdChhWyJpbWFnZV9pZCJdLCBbXSkuYXBwZW5kKGEpCiAgICAgICAgc2VsZi5yZWNvcmRzID0gW10KICAgICAgICBmb3IgaW1hZ2VfaWQsIGluZm8gaW4gc2VsZi5pbWFnZXMuaXRlbXMoKToKICAgICAgICAgICAgc2VsZi5yZWNvcmRzLmFwcGVuZCh7CiAgICAgICAgICAgICAgICAiaW1hZ2VfaWQiOiBpbWFnZV9pZCwgImZpbGVfbmFtZSI6IGluZm9bImZpbGVfbmFtZSJdLAogICAgICAgICAgICAgICAgIndpZHRoIjogaW5mb1sid2lkdGgiXSwgImhlaWdodCI6IGluZm9bImhlaWdodCJdLAogICAgICAgICAgICAgICAgImFubm90YXRpb25zIjogYW5ucy5nZXQoaW1hZ2VfaWQsIFtdKQogICAgICAgICAgICB9KQoKICAgIGRlZiBfX2xlbl9fKHNlbGYpOiByZXR1cm4gbGVuKHNlbGYucmVjb3JkcykKCiAgICBkZWYgX19nZXRpdGVtX18oc2VsZiwgaWR4KToKICAgICAgICByID0gc2VsZi5yZWNvcmRzW2lkeF0KICAgICAgICBpbWFnZSA9IEltYWdlLm9wZW4oc2VsZi5pbWFnZV9kaXIgLyByWyJmaWxlX25hbWUiXSkuY29udmVydCgiUkdCIikKICAgICAgICBhbm5zID0gW10KICAgICAgICBmb3IgYSBpbiByWyJhbm5vdGF0aW9ucyJdOgogICAgICAgICAgICB4LHksdyxoID0gYVsiYmJveCJdCiAgICAgICAgICAgIGlmIHcgPD0gMCBvciBoIDw9IDA6IGNvbnRpbnVlCiAgICAgICAgICAgIGFubnMuYXBwZW5kKHsKICAgICAgICAgICAgICAgICJpZCI6IGFbImlkIl0sICJpbWFnZV9pZCI6IGludChpZHgpLAogICAgICAgICAgICAgICAgImNhdGVnb3J5X2lkIjogc2VsZi5jYXRlZ29yeV9pZF90b19sYWJlbFthWyJjYXRlZ29yeV9pZCJdXSwKICAgICAgICAgICAgICAgICJiYm94IjogW3gseSx3LGhdLCAiYXJlYSI6IGZsb2F0KGEuZ2V0KCJhcmVhIix3KmgpKSwKICAgICAgICAgICAgICAgICJpc2Nyb3dkIjogMAogICAgICAgICAgICB9KQogICAgICAgIGVuY29kZWQgPSBzZWxmLnByb2Nlc3NvcigKICAgICAgICAgICAgaW1hZ2VzPWltYWdlLAogICAgICAgICAgICBhbm5vdGF0aW9ucz17ImltYWdlX2lkIjogaW50KGlkeCksICJhbm5vdGF0aW9ucyI6IGFubnN9LAogICAgICAgICAgICByZXR1cm5fdGVuc29ycz0icHQiCiAgICAgICAgKQogICAgICAgIGVuY29kZWRbInBpeGVsX3ZhbHVlcyJdID0gZW5jb2RlZFsicGl4ZWxfdmFsdWVzIl0uc3F1ZWV6ZSgwKQogICAgICAgIGlmICJwaXhlbF9tYXNrIiBpbiBlbmNvZGVkOgogICAgICAgICAgICBlbmNvZGVkWyJwaXhlbF9tYXNrIl0gPSBlbmNvZGVkWyJwaXhlbF9tYXNrIl0uc3F1ZWV6ZSgwKQogICAgICAgIGVuY29kZWRbImxhYmVscyJdID0gZW5jb2RlZFsibGFiZWxzIl1bMF0KICAgICAgICByZXR1cm4gZW5jb2RlZAoKZGVmIGNvbGxhdGVfZm4oYmF0Y2gpOgogICAgb3V0ID0geyJwaXhlbF92YWx1ZXMiOiB0b3JjaC5zdGFjayhbeFsicGl4ZWxfdmFsdWVzIl0gZm9yIHggaW4gYmF0Y2hdKSwKICAgICAgICAgICAibGFiZWxzIjogW3hbImxhYmVscyJdIGZvciB4IGluIGJhdGNoXX0KICAgIGlmICJwaXhlbF9tYXNrIiBpbiBiYXRjaFswXToKICAgICAgICBvdXRbInBpeGVsX21hc2siXSA9IHRvcmNoLnN0YWNrKFt4WyJwaXhlbF9tYXNrIl0gZm9yIHggaW4gYmF0Y2hdKQogICAgcmV0dXJuIG91dAoKZGVmIG1vdmVfdG9fZGV2aWNlKG9iaiwgZGV2aWNlKToKICAgIGlmIHRvcmNoLmlzX3RlbnNvcihvYmopOgogICAgICAgIHJldHVybiBvYmoudG8oZGV2aWNlKQogICAgaWYgaXNpbnN0YW5jZShvYmosIGRpY3QpOgogICAgICAgIHJldHVybiB7azogbW92ZV90b19kZXZpY2UodiwgZGV2aWNlKSBmb3IgaywgdiBpbiBvYmouaXRlbXMoKX0KICAgIGlmIGlzaW5zdGFuY2Uob2JqLCBsaXN0KToKICAgICAgICByZXR1cm4gW21vdmVfdG9fZGV2aWNlKHYsIGRldmljZSkgZm9yIHYgaW4gb2JqXQogICAgaWYgaXNpbnN0YW5jZShvYmosIHR1cGxlKToKICAgICAgICByZXR1cm4gdHVwbGUobW92ZV90b19kZXZpY2UodiwgZGV2aWNlKSBmb3IgdiBpbiBvYmopCiAgICByZXR1cm4gb2JqCgpkZWYgZXZhbHVhdGUobW9kZWwsIGxvYWRlciwgZGV2aWNlKToKICAgIG1vZGVsLmV2YWwoKTsgdG90YWw9MDsgbj0wCiAgICB3aXRoIHRvcmNoLm5vX2dyYWQoKToKICAgICAgICBmb3IgYmF0Y2ggaW4gbG9hZGVyOgogICAgICAgICAgICBiYXRjaD1tb3ZlX3RvX2RldmljZShiYXRjaCwgZGV2aWNlKQogICAgICAgICAgICB0b3RhbCArPSBmbG9hdChtb2RlbCgqKmJhdGNoKS5sb3NzLml0ZW0oKSk7IG4gKz0gMQogICAgbW9kZWwudHJhaW4oKQogICAgcmV0dXJuIHRvdGFsL21heChuLDEpCgpkZWYgbWFpbigpOgogICAgcD1hcmdwYXJzZS5Bcmd1bWVudFBhcnNlcigpCiAgICBwLmFkZF9hcmd1bWVudCgiLS10cmFpbi1kaXIiLHJlcXVpcmVkPVRydWUpOyBwLmFkZF9hcmd1bWVudCgiLS12YWwtZGlyIixyZXF1aXJlZD1UcnVlKQogICAgcC5hZGRfYXJndW1lbnQoIi0tY2xhc3NlcyIscmVxdWlyZWQ9VHJ1ZSk7IHAuYWRkX2FyZ3VtZW50KCItLW91dHB1dC1kaXIiLGRlZmF1bHQ9Im1vZGVsIikKICAgIHAuYWRkX2FyZ3VtZW50KCItLWVwb2NocyIsdHlwZT1pbnQsZGVmYXVsdD0zMCk7IHAuYWRkX2FyZ3VtZW50KCItLWJhdGNoLXNpemUiLHR5cGU9aW50LGRlZmF1bHQ9MikKICAgIHAuYWRkX2FyZ3VtZW50KCItLWxlYXJuaW5nLXJhdGUiLHR5cGU9ZmxvYXQsZGVmYXVsdD0xZS01KTsgcC5hZGRfYXJndW1lbnQoIi0td2VpZ2h0LWRlY2F5Iix0eXBlPWZsb2F0LGRlZmF1bHQ9MWUtNCkKICAgIHAuYWRkX2FyZ3VtZW50KCItLW51bS13b3JrZXJzIix0eXBlPWludCxkZWZhdWx0PTIpCiAgICBhPXAucGFyc2VfYXJncygpCgogICAgY2xhc3Nlcz1sb2FkX2NsYXNzZXMoYS5jbGFzc2VzKQogICAgaWQybGFiZWw9e2k6biBmb3IgaSxuIGluIGVudW1lcmF0ZShjbGFzc2VzKX0KICAgIGxhYmVsMmlkPXtuOmkgZm9yIGksbiBpbiBlbnVtZXJhdGUoY2xhc3Nlcyl9CgogICAgcHJvYz1SVERldHJJbWFnZVByb2Nlc3Nvci5mcm9tX3ByZXRyYWluZWQoQkFTRV9NT0RFTCkKICAgIHRyYWluPUNPQ09EZXRlY3Rpb25EYXRhc2V0KFBhdGgoYS50cmFpbl9kaXIpLyJpbWFnZXMiLFBhdGgoYS50cmFpbl9kaXIpLyJhbm5vdGF0aW9ucy5qc29uIixwcm9jKQogICAgdmFsPUNPQ09EZXRlY3Rpb25EYXRhc2V0KFBhdGgoYS52YWxfZGlyKS8iaW1hZ2VzIixQYXRoKGEudmFsX2RpcikvImFubm90YXRpb25zLmpzb24iLHByb2MpCgogICAgaWYgbGVuKHRyYWluKT09MCBvciBsZW4odmFsKT09MDoKICAgICAgICByYWlzZSBWYWx1ZUVycm9yKCJUcmFpbmluZyBhbmQgdmFsaWRhdGlvbiBkYXRhc2V0cyBtdXN0IGNvbnRhaW4gYXQgbGVhc3Qgb25lIGltYWdlLiIpCiAgICBpZiBsZW4odHJhaW4uY2F0ZWdvcnlfaWRfdG9fbGFiZWwpIT1sZW4oY2xhc3Nlcykgb3IgbGVuKHZhbC5jYXRlZ29yeV9pZF90b19sYWJlbCkhPWxlbihjbGFzc2VzKToKICAgICAgICByYWlzZSBWYWx1ZUVycm9yKCJDT0NPIGNhdGVnb3JpZXMgZG8gbm90IG1hdGNoIGNsYXNzZXMudHh0LiBSZWJ1aWxkIHRoZSBkYXRhc2V0IGFmdGVyIHNhdmluZyB0aGUgY2xhc3Nlcy4iKQoKICAgIG1vZGVsPVJURGV0ckZvck9iamVjdERldGVjdGlvbi5mcm9tX3ByZXRyYWluZWQoCiAgICAgICAgQkFTRV9NT0RFTCxudW1fbGFiZWxzPWxlbihjbGFzc2VzKSxpZDJsYWJlbD1pZDJsYWJlbCxsYWJlbDJpZD1sYWJlbDJpZCwKICAgICAgICBpZ25vcmVfbWlzbWF0Y2hlZF9zaXplcz1UcnVlCiAgICApCiAgICBkZXZpY2U9dG9yY2guZGV2aWNlKCJjdWRhIiBpZiB0b3JjaC5jdWRhLmlzX2F2YWlsYWJsZSgpIGVsc2UgImNwdSIpCiAgICBtb2RlbC50byhkZXZpY2UpCgogICAgdHI9RGF0YUxvYWRlcih0cmFpbixiYXRjaF9zaXplPWEuYmF0Y2hfc2l6ZSxzaHVmZmxlPVRydWUsbnVtX3dvcmtlcnM9MCxjb2xsYXRlX2ZuPWNvbGxhdGVfZm4pCiAgICB2YT1EYXRhTG9hZGVyKHZhbCxiYXRjaF9zaXplPWEuYmF0Y2hfc2l6ZSxzaHVmZmxlPUZhbHNlLG51bV93b3JrZXJzPTAsY29sbGF0ZV9mbj1jb2xsYXRlX2ZuKQogICAgb3B0PXRvcmNoLm9wdGltLkFkYW1XKG1vZGVsLnBhcmFtZXRlcnMoKSxscj1hLmxlYXJuaW5nX3JhdGUsd2VpZ2h0X2RlY2F5PWEud2VpZ2h0X2RlY2F5KQoKICAgIG91dGRpcj1QYXRoKGEub3V0cHV0X2Rpcik7IG91dGRpci5ta2RpcihwYXJlbnRzPVRydWUsZXhpc3Rfb2s9VHJ1ZSkKICAgIGJlc3Q9ZmxvYXQoImluZiIpCgogICAgZm9yIGVwb2NoIGluIHJhbmdlKGEuZXBvY2hzKToKICAgICAgICBtb2RlbC50cmFpbigpOyBydW5uaW5nPTAKICAgICAgICBiYXI9dHFkbSh0cixkZXNjPWYiZXBvY2gge2Vwb2NoKzF9L3thLmVwb2Noc30iKQogICAgICAgIGZvciBzdGVwLGJhdGNoIGluIGVudW1lcmF0ZShiYXIpOgogICAgICAgICAgICBiYXRjaD1tb3ZlX3RvX2RldmljZShiYXRjaCwgZGV2aWNlKQogICAgICAgICAgICAjIFJULURFVFIncyBsb3NzIG1hdGNoZXIgdXNlcyBuZXN0ZWQgdGFyZ2V0IHRlbnNvcnMgKGJveGVzL2NsYXNzZXMpLgogICAgICAgICAgICAjIE1vdmUgZXZlcnkgdGVuc29yIGluIGxhYmVscyB0byB0aGUgc2FtZSBkZXZpY2UgYXMgdGhlIG1vZGVsLgogICAgICAgICAgICBpZiAibGFiZWxzIiBpbiBiYXRjaDoKICAgICAgICAgICAgICAgIGJhdGNoWyJsYWJlbHMiXSA9IG1vdmVfdG9fZGV2aWNlKGJhdGNoWyJsYWJlbHMiXSwgZGV2aWNlKQogICAgICAgICAgICBsb3NzPW1vZGVsKCoqYmF0Y2gpLmxvc3MKICAgICAgICAgICAgbG9zcy5iYWNrd2FyZCgpOyBvcHQuc3RlcCgpOyBvcHQuemVyb19ncmFkKHNldF90b19ub25lPVRydWUpCiAgICAgICAgICAgIHJ1bm5pbmcgKz0gZmxvYXQobG9zcy5pdGVtKCkpCiAgICAgICAgICAgIGJhci5zZXRfcG9zdGZpeChsb3NzPWYie3J1bm5pbmcvKHN0ZXArMSk6LjRmfSIpCiAgICAgICAgdmw9ZXZhbHVhdGUobW9kZWwsdmEsZGV2aWNlKQogICAgICAgIHByaW50KGYidmFsaWRhdGlvbl9sb3NzPXt2bDouNGZ9IikKICAgICAgICBpZiB2bDxiZXN0OgogICAgICAgICAgICBiZXN0PXZsCiAgICAgICAgICAgIG1vZGVsLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpCiAgICAgICAgICAgIHByb2Muc2F2ZV9wcmV0cmFpbmVkKG91dGRpcikKICAgICAgICAgICAgKG91dGRpci8iY2xhc3Nlcy5qc29uIikud3JpdGVfdGV4dChqc29uLmR1bXBzKHsiaWQybGFiZWwiOmlkMmxhYmVsLCJsYWJlbDJpZCI6bGFiZWwyaWR9LGluZGVudD0yKSkKICAgIG1vZGVsLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpOyBwcm9jLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpCgppZiBfX25hbWVfXz09Il9fbWFpbl9fIjogbWFpbigpCg=="
|
| 30 |
TRAINING_DIR = ROOT / "training"
|
| 31 |
TRAIN_SCRIPT = TRAINING_DIR / "train.py"
|
| 32 |
if not TRAIN_SCRIPT.exists():
|
|
|
|
| 26 |
ROOT = Path(__file__).resolve().parent
|
| 27 |
|
| 28 |
# Self-heal the training script if a deployment omitted the training/ directory.
|
| 29 |
+
_TRAINING_SCRIPT_B64 = "aW1wb3J0IGFyZ3BhcnNlCmltcG9ydCBqc29uCmZyb20gcGF0aGxpYiBpbXBvcnQgUGF0aAoKaW1wb3J0IHRvcmNoCmZyb20gUElMIGltcG9ydCBJbWFnZQpmcm9tIHRvcmNoLnV0aWxzLmRhdGEgaW1wb3J0IERhdGFzZXQsIERhdGFMb2FkZXIKZnJvbSB0cWRtIGltcG9ydCB0cWRtCmZyb20gdHJhbnNmb3JtZXJzIGltcG9ydCBSVERldHJJbWFnZVByb2Nlc3NvciwgUlREZXRyRm9yT2JqZWN0RGV0ZWN0aW9uCgpCQVNFX01PREVMID0gIlBla2luZ1UvcnRkZXRyX3I1MHZkIgoKZGVmIGxvYWRfY2xhc3NlcyhwYXRoKToKICAgIHJldHVybiBbeC5zdHJpcCgpIGZvciB4IGluIFBhdGgocGF0aCkucmVhZF90ZXh0KCkuc3BsaXRsaW5lcygpIGlmIHguc3RyaXAoKV0KCmNsYXNzIENPQ09EZXRlY3Rpb25EYXRhc2V0KERhdGFzZXQpOgogICAgZGVmIF9faW5pdF9fKHNlbGYsIGltYWdlX2RpciwgYW5ub3RhdGlvbl9maWxlLCBwcm9jZXNzb3IpOgogICAgICAgIHNlbGYuaW1hZ2VfZGlyID0gUGF0aChpbWFnZV9kaXIpCiAgICAgICAgc2VsZi5wcm9jZXNzb3IgPSBwcm9jZXNzb3IKICAgICAgICBjb2NvID0ganNvbi5sb2FkcyhQYXRoKGFubm90YXRpb25fZmlsZSkucmVhZF90ZXh0KCkpCiAgICAgICAgc2VsZi5pbWFnZXMgPSB7eFsiaWQiXTogeCBmb3IgeCBpbiBjb2NvWyJpbWFnZXMiXX0KICAgICAgICBjYXRzID0gc29ydGVkKGNvY29bImNhdGVnb3JpZXMiXSwga2V5PWxhbWJkYSB4OiB4WyJpZCJdKQogICAgICAgIHNlbGYuY2F0ZWdvcnlfaWRfdG9fbGFiZWwgPSB7Y1siaWQiXTogaSBmb3IgaSxjIGluIGVudW1lcmF0ZShjYXRzKX0KICAgICAgICBhbm5zID0ge30KICAgICAgICBmb3IgYSBpbiBjb2NvWyJhbm5vdGF0aW9ucyJdOgogICAgICAgICAgICBpZiBub3QgYS5nZXQoImlzY3Jvd2QiLCAwKToKICAgICAgICAgICAgICAgIGFubnMuc2V0ZGVmYXVsdChhWyJpbWFnZV9pZCJdLCBbXSkuYXBwZW5kKGEpCiAgICAgICAgc2VsZi5yZWNvcmRzID0gW10KICAgICAgICBmb3IgaW1hZ2VfaWQsIGluZm8gaW4gc2VsZi5pbWFnZXMuaXRlbXMoKToKICAgICAgICAgICAgc2VsZi5yZWNvcmRzLmFwcGVuZCh7CiAgICAgICAgICAgICAgICAiaW1hZ2VfaWQiOiBpbWFnZV9pZCwgImZpbGVfbmFtZSI6IGluZm9bImZpbGVfbmFtZSJdLAogICAgICAgICAgICAgICAgIndpZHRoIjogaW5mb1sid2lkdGgiXSwgImhlaWdodCI6IGluZm9bImhlaWdodCJdLAogICAgICAgICAgICAgICAgImFubm90YXRpb25zIjogYW5ucy5nZXQoaW1hZ2VfaWQsIFtdKQogICAgICAgICAgICB9KQoKICAgIGRlZiBfX2xlbl9fKHNlbGYpOiByZXR1cm4gbGVuKHNlbGYucmVjb3JkcykKCiAgICBkZWYgX19nZXRpdGVtX18oc2VsZiwgaWR4KToKICAgICAgICByID0gc2VsZi5yZWNvcmRzW2lkeF0KICAgICAgICBpbWFnZSA9IEltYWdlLm9wZW4oc2VsZi5pbWFnZV9kaXIgLyByWyJmaWxlX25hbWUiXSkuY29udmVydCgiUkdCIikKICAgICAgICBhbm5zID0gW10KICAgICAgICBmb3IgYSBpbiByWyJhbm5vdGF0aW9ucyJdOgogICAgICAgICAgICB4LHksdyxoID0gYVsiYmJveCJdCiAgICAgICAgICAgIGlmIHcgPD0gMCBvciBoIDw9IDA6IGNvbnRpbnVlCiAgICAgICAgICAgIGFubnMuYXBwZW5kKHsKICAgICAgICAgICAgICAgICJpZCI6IGFbImlkIl0sICJpbWFnZV9pZCI6IGludChpZHgpLAogICAgICAgICAgICAgICAgImNhdGVnb3J5X2lkIjogc2VsZi5jYXRlZ29yeV9pZF90b19sYWJlbFthWyJjYXRlZ29yeV9pZCJdXSwKICAgICAgICAgICAgICAgICJiYm94IjogW3gseSx3LGhdLCAiYXJlYSI6IGZsb2F0KGEuZ2V0KCJhcmVhIix3KmgpKSwKICAgICAgICAgICAgICAgICJpc2Nyb3dkIjogMAogICAgICAgICAgICB9KQogICAgICAgIGVuY29kZWQgPSBzZWxmLnByb2Nlc3NvcigKICAgICAgICAgICAgaW1hZ2VzPWltYWdlLAogICAgICAgICAgICBhbm5vdGF0aW9ucz17ImltYWdlX2lkIjogaW50KGlkeCksICJhbm5vdGF0aW9ucyI6IGFubnN9LAogICAgICAgICAgICByZXR1cm5fdGVuc29ycz0icHQiCiAgICAgICAgKQogICAgICAgIGVuY29kZWRbInBpeGVsX3ZhbHVlcyJdID0gZW5jb2RlZFsicGl4ZWxfdmFsdWVzIl0uc3F1ZWV6ZSgwKQogICAgICAgIGlmICJwaXhlbF9tYXNrIiBpbiBlbmNvZGVkOgogICAgICAgICAgICBlbmNvZGVkWyJwaXhlbF9tYXNrIl0gPSBlbmNvZGVkWyJwaXhlbF9tYXNrIl0uc3F1ZWV6ZSgwKQogICAgICAgIGVuY29kZWRbImxhYmVscyJdID0gZW5jb2RlZFsibGFiZWxzIl1bMF0KICAgICAgICByZXR1cm4gZW5jb2RlZAoKZGVmIGNvbGxhdGVfZm4oYmF0Y2gpOgogICAgb3V0ID0geyJwaXhlbF92YWx1ZXMiOiB0b3JjaC5zdGFjayhbeFsicGl4ZWxfdmFsdWVzIl0gZm9yIHggaW4gYmF0Y2hdKSwKICAgICAgICAgICAibGFiZWxzIjogW3hbImxhYmVscyJdIGZvciB4IGluIGJhdGNoXX0KICAgIGlmICJwaXhlbF9tYXNrIiBpbiBiYXRjaFswXToKICAgICAgICBvdXRbInBpeGVsX21hc2siXSA9IHRvcmNoLnN0YWNrKFt4WyJwaXhlbF9tYXNrIl0gZm9yIHggaW4gYmF0Y2hdKQogICAgcmV0dXJuIG91dAoKZGVmIG1vdmVfdG9fZGV2aWNlKG9iaiwgZGV2aWNlKToKICAgIGlmIHRvcmNoLmlzX3RlbnNvcihvYmopOgogICAgICAgIHJldHVybiBvYmoudG8oZGV2aWNlKQogICAgaWYgaXNpbnN0YW5jZShvYmosIGRpY3QpOgogICAgICAgIHJldHVybiB7azogbW92ZV90b19kZXZpY2UodiwgZGV2aWNlKSBmb3IgaywgdiBpbiBvYmouaXRlbXMoKX0KICAgIGlmIGlzaW5zdGFuY2Uob2JqLCBsaXN0KToKICAgICAgICByZXR1cm4gW21vdmVfdG9fZGV2aWNlKHYsIGRldmljZSkgZm9yIHYgaW4gb2JqXQogICAgaWYgaXNpbnN0YW5jZShvYmosIHR1cGxlKToKICAgICAgICByZXR1cm4gdHVwbGUobW92ZV90b19kZXZpY2UodiwgZGV2aWNlKSBmb3IgdiBpbiBvYmopCiAgICByZXR1cm4gb2JqCgpkZWYgZXZhbHVhdGUobW9kZWwsIGxvYWRlciwgZGV2aWNlKToKICAgIG1vZGVsLmV2YWwoKTsgdG90YWw9MDsgbj0wCiAgICB3aXRoIHRvcmNoLm5vX2dyYWQoKToKICAgICAgICBmb3IgYmF0Y2ggaW4gbG9hZGVyOgogICAgICAgICAgICBiYXRjaD1tb3ZlX3RvX2RldmljZShiYXRjaCwgZGV2aWNlKQogICAgICAgICAgICB0b3RhbCArPSBmbG9hdChtb2RlbCgqKmJhdGNoKS5sb3NzLml0ZW0oKSk7IG4gKz0gMQogICAgbW9kZWwudHJhaW4oKQogICAgcmV0dXJuIHRvdGFsL21heChuLDEpCgoKZGVmIHBhdGNoX3J0ZGV0cl9kZW5vaXNpbmdfZGV2aWNlKCk6CiAgICAiIiJXb3JrIGFyb3VuZCBSVC1ERVRSIGRlbm9pc2luZyBjb2RlIGNyZWF0aW5nIENQVSBpbmRleCB0ZW5zb3JzIG9uIHNvbWUgVHJhbnNmb3JtZXJzIHJlbGVhc2VzLiIiIgogICAgdHJ5OgogICAgICAgIGltcG9ydCB0cmFuc2Zvcm1lcnMubW9kZWxzLnJ0X2RldHIubW9kZWxpbmdfcnRfZGV0ciBhcyBydGRldHJfbW9kCiAgICAgICAgb3JpZ2luYWwgPSBydGRldHJfbW9kLmdldF9jb250cmFzdGl2ZV9kZW5vaXNpbmdfdHJhaW5pbmdfZ3JvdXAKICAgICAgICBpZiBnZXRhdHRyKG9yaWdpbmFsLCAiX2ljZWNyZWFtX2RldmljZV9wYXRjaCIsIEZhbHNlKToKICAgICAgICAgICAgcmV0dXJuCgogICAgICAgIGRlZiB3cmFwcGVkKHRhcmdldHMsICphcmdzLCAqKmt3YXJncyk6CiAgICAgICAgICAgICMgY2xhc3NfZW1iZWQgaXMgdGhlIDR0aCBwb3NpdGlvbmFsIGFyZ3VtZW50IGluIHRoZSBzdXBwb3J0ZWQgUlQtREVUUiB2ZXJzaW9ucy4KICAgICAgICAgICAgY2xhc3NfZW1iZWQgPSBhcmdzWzJdIGlmIGxlbihhcmdzKSA+PSAzIGVsc2Uga3dhcmdzLmdldCgiY2xhc3NfZW1iZWQiKQogICAgICAgICAgICB0cnk6CiAgICAgICAgICAgICAgICBkZXZpY2UgPSBuZXh0KGNsYXNzX2VtYmVkLnBhcmFtZXRlcnMoKSkuZGV2aWNlCiAgICAgICAgICAgIGV4Y2VwdCBFeGNlcHRpb246CiAgICAgICAgICAgICAgICBkZXZpY2UgPSBOb25lCiAgICAgICAgICAgIGlmIGRldmljZSBpcyBub3QgTm9uZToKICAgICAgICAgICAgICAgIGZvciB0YXJnZXQgaW4gdGFyZ2V0czoKICAgICAgICAgICAgICAgICAgICBpZiBpc2luc3RhbmNlKHRhcmdldCwgZGljdCk6CiAgICAgICAgICAgICAgICAgICAgICAgIGZvciBrZXkgaW4gKCJjbGFzc19sYWJlbHMiLCAiYm94ZXMiKToKICAgICAgICAgICAgICAgICAgICAgICAgICAgIHZhbHVlID0gdGFyZ2V0LmdldChrZXkpCiAgICAgICAgICAgICAgICAgICAgICAgICAgICBpZiB0b3JjaC5pc190ZW5zb3IodmFsdWUpIGFuZCB2YWx1ZS5kZXZpY2UgIT0gZGV2aWNlOgogICAgICAgICAgICAgICAgICAgICAgICAgICAgICAgIHRhcmdldFtrZXldID0gdmFsdWUudG8oZGV2aWNlKQogICAgICAgICAgICByZXR1cm4gb3JpZ2luYWwodGFyZ2V0cywgKmFyZ3MsICoqa3dhcmdzKQoKICAgICAgICB3cmFwcGVkLl9pY2VjcmVhbV9kZXZpY2VfcGF0Y2ggPSBUcnVlCiAgICAgICAgcnRkZXRyX21vZC5nZXRfY29udHJhc3RpdmVfZGVub2lzaW5nX3RyYWluaW5nX2dyb3VwID0gd3JhcHBlZAogICAgZXhjZXB0IEV4Y2VwdGlvbiBhcyBleGM6CiAgICAgICAgcHJpbnQoZiJXYXJuaW5nOiBSVC1ERVRSIGRlbm9pc2luZyBkZXZpY2UgcGF0Y2ggd2FzIG5vdCBpbnN0YWxsZWQ6IHtleGN9IikKCmRlZiBtYWluKCk6CiAgICBwPWFyZ3BhcnNlLkFyZ3VtZW50UGFyc2VyKCkKICAgIHAuYWRkX2FyZ3VtZW50KCItLXRyYWluLWRpciIscmVxdWlyZWQ9VHJ1ZSk7IHAuYWRkX2FyZ3VtZW50KCItLXZhbC1kaXIiLHJlcXVpcmVkPVRydWUpCiAgICBwLmFkZF9hcmd1bWVudCgiLS1jbGFzc2VzIixyZXF1aXJlZD1UcnVlKTsgcC5hZGRfYXJndW1lbnQoIi0tb3V0cHV0LWRpciIsZGVmYXVsdD0ibW9kZWwiKQogICAgcC5hZGRfYXJndW1lbnQoIi0tZXBvY2hzIix0eXBlPWludCxkZWZhdWx0PTMwKTsgcC5hZGRfYXJndW1lbnQoIi0tYmF0Y2gtc2l6ZSIsdHlwZT1pbnQsZGVmYXVsdD0yKQogICAgcC5hZGRfYXJndW1lbnQoIi0tbGVhcm5pbmctcmF0ZSIsdHlwZT1mbG9hdCxkZWZhdWx0PTFlLTUpOyBwLmFkZF9hcmd1bWVudCgiLS13ZWlnaHQtZGVjYXkiLHR5cGU9ZmxvYXQsZGVmYXVsdD0xZS00KQogICAgcC5hZGRfYXJndW1lbnQoIi0tbnVtLXdvcmtlcnMiLHR5cGU9aW50LGRlZmF1bHQ9MikKICAgIGE9cC5wYXJzZV9hcmdzKCkKCiAgICBjbGFzc2VzPWxvYWRfY2xhc3NlcyhhLmNsYXNzZXMpCiAgICBpZDJsYWJlbD17aTpuIGZvciBpLG4gaW4gZW51bWVyYXRlKGNsYXNzZXMpfQogICAgbGFiZWwyaWQ9e246aSBmb3IgaSxuIGluIGVudW1lcmF0ZShjbGFzc2VzKX0KCiAgICBwcm9jPVJURGV0ckltYWdlUHJvY2Vzc29yLmZyb21fcHJldHJhaW5lZChCQVNFX01PREVMKQogICAgdHJhaW49Q09DT0RldGVjdGlvbkRhdGFzZXQoUGF0aChhLnRyYWluX2RpcikvImltYWdlcyIsUGF0aChhLnRyYWluX2RpcikvImFubm90YXRpb25zLmpzb24iLHByb2MpCiAgICB2YWw9Q09DT0RldGVjdGlvbkRhdGFzZXQoUGF0aChhLnZhbF9kaXIpLyJpbWFnZXMiLFBhdGgoYS52YWxfZGlyKS8iYW5ub3RhdGlvbnMuanNvbiIscHJvYykKCiAgICBpZiBsZW4odHJhaW4pPT0wIG9yIGxlbih2YWwpPT0wOgogICAgICAgIHJhaXNlIFZhbHVlRXJyb3IoIlRyYWluaW5nIGFuZCB2YWxpZGF0aW9uIGRhdGFzZXRzIG11c3QgY29udGFpbiBhdCBsZWFzdCBvbmUgaW1hZ2UuIikKICAgIGlmIGxlbih0cmFpbi5jYXRlZ29yeV9pZF90b19sYWJlbCkhPWxlbihjbGFzc2VzKSBvciBsZW4odmFsLmNhdGVnb3J5X2lkX3RvX2xhYmVsKSE9bGVuKGNsYXNzZXMpOgogICAgICAgIHJhaXNlIFZhbHVlRXJyb3IoIkNPQ08gY2F0ZWdvcmllcyBkbyBub3QgbWF0Y2ggY2xhc3Nlcy50eHQuIFJlYnVpbGQgdGhlIGRhdGFzZXQgYWZ0ZXIgc2F2aW5nIHRoZSBjbGFzc2VzLiIpCgogICAgbW9kZWw9UlREZXRyRm9yT2JqZWN0RGV0ZWN0aW9uLmZyb21fcHJldHJhaW5lZCgKICAgICAgICBCQVNFX01PREVMLG51bV9sYWJlbHM9bGVuKGNsYXNzZXMpLGlkMmxhYmVsPWlkMmxhYmVsLGxhYmVsMmlkPWxhYmVsMmlkLAogICAgICAgIGlnbm9yZV9taXNtYXRjaGVkX3NpemVzPVRydWUKICAgICkKICAgIGRldmljZT10b3JjaC5kZXZpY2UoImN1ZGEiIGlmIHRvcmNoLmN1ZGEuaXNfYXZhaWxhYmxlKCkgZWxzZSAiY3B1IikKICAgIG1vZGVsLnRvKGRldmljZSkKICAgIHBhdGNoX3J0ZGV0cl9kZW5vaXNpbmdfZGV2aWNlKCkKCiAgICB0cj1EYXRhTG9hZGVyKHRyYWluLGJhdGNoX3NpemU9YS5iYXRjaF9zaXplLHNodWZmbGU9VHJ1ZSxudW1fd29ya2Vycz0wLGNvbGxhdGVfZm49Y29sbGF0ZV9mbikKICAgIHZhPURhdGFMb2FkZXIodmFsLGJhdGNoX3NpemU9YS5iYXRjaF9zaXplLHNodWZmbGU9RmFsc2UsbnVtX3dvcmtlcnM9MCxjb2xsYXRlX2ZuPWNvbGxhdGVfZm4pCiAgICBvcHQ9dG9yY2gub3B0aW0uQWRhbVcobW9kZWwucGFyYW1ldGVycygpLGxyPWEubGVhcm5pbmdfcmF0ZSx3ZWlnaHRfZGVjYXk9YS53ZWlnaHRfZGVjYXkpCgogICAgb3V0ZGlyPVBhdGgoYS5vdXRwdXRfZGlyKTsgb3V0ZGlyLm1rZGlyKHBhcmVudHM9VHJ1ZSxleGlzdF9vaz1UcnVlKQogICAgYmVzdD1mbG9hdCgiaW5mIikKCiAgICBmb3IgZXBvY2ggaW4gcmFuZ2UoYS5lcG9jaHMpOgogICAgICAgIG1vZGVsLnRyYWluKCk7IHJ1bm5pbmc9MAogICAgICAgIGJhcj10cWRtKHRyLGRlc2M9ZiJlcG9jaCB7ZXBvY2grMX0ve2EuZXBvY2hzfSIpCiAgICAgICAgZm9yIHN0ZXAsYmF0Y2ggaW4gZW51bWVyYXRlKGJhcik6CiAgICAgICAgICAgIGJhdGNoPW1vdmVfdG9fZGV2aWNlKGJhdGNoLCBkZXZpY2UpCiAgICAgICAgICAgICMgUlQtREVUUidzIGxvc3MgbWF0Y2hlciB1c2VzIG5lc3RlZCB0YXJnZXQgdGVuc29ycyAoYm94ZXMvY2xhc3NlcykuCiAgICAgICAgICAgICMgTW92ZSBldmVyeSB0ZW5zb3IgaW4gbGFiZWxzIHRvIHRoZSBzYW1lIGRldmljZSBhcyB0aGUgbW9kZWwuCiAgICAgICAgICAgIGlmICJsYWJlbHMiIGluIGJhdGNoOgogICAgICAgICAgICAgICAgIyBSVC1ERVRSIGV4cGVjdHMgZXZlcnkgbmVzdGVkIHRhcmdldCB0ZW5zb3Igb24gdGhlIHNhbWUgZGV2aWNlIGFzIHRoZSBtb2RlbC4KICAgICAgICAgICAgICAgIGZvciB0YXJnZXQgaW4gYmF0Y2hbImxhYmVscyJdOgogICAgICAgICAgICAgICAgICAgIGlmIGlzaW5zdGFuY2UodGFyZ2V0LCBkaWN0KToKICAgICAgICAgICAgICAgICAgICAgICAgZm9yIGtleSwgdmFsdWUgaW4gbGlzdCh0YXJnZXQuaXRlbXMoKSk6CiAgICAgICAgICAgICAgICAgICAgICAgICAgICBpZiB0b3JjaC5pc190ZW5zb3IodmFsdWUpOgogICAgICAgICAgICAgICAgICAgICAgICAgICAgICAgIHRhcmdldFtrZXldID0gdmFsdWUudG8oZGV2aWNlKQogICAgICAgICAgICBsb3NzPW1vZGVsKCoqYmF0Y2gpLmxvc3MKICAgICAgICAgICAgbG9zcy5iYWNrd2FyZCgpOyBvcHQuc3RlcCgpOyBvcHQuemVyb19ncmFkKHNldF90b19ub25lPVRydWUpCiAgICAgICAgICAgIHJ1bm5pbmcgKz0gZmxvYXQobG9zcy5pdGVtKCkpCiAgICAgICAgICAgIGJhci5zZXRfcG9zdGZpeChsb3NzPWYie3J1bm5pbmcvKHN0ZXArMSk6LjRmfSIpCiAgICAgICAgdmw9ZXZhbHVhdGUobW9kZWwsdmEsZGV2aWNlKQogICAgICAgIHByaW50KGYidmFsaWRhdGlvbl9sb3NzPXt2bDouNGZ9IikKICAgICAgICBpZiB2bDxiZXN0OgogICAgICAgICAgICBiZXN0PXZsCiAgICAgICAgICAgIG1vZGVsLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpCiAgICAgICAgICAgIHByb2Muc2F2ZV9wcmV0cmFpbmVkKG91dGRpcikKICAgICAgICAgICAgKG91dGRpci8iY2xhc3Nlcy5qc29uIikud3JpdGVfdGV4dChqc29uLmR1bXBzKHsiaWQybGFiZWwiOmlkMmxhYmVsLCJsYWJlbDJpZCI6bGFiZWwyaWR9LGluZGVudD0yKSkKICAgIG1vZGVsLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpOyBwcm9jLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpCgppZiBfX25hbWVfXz09Il9fbWFpbl9fIjogbWFpbigpCg=="
|
| 30 |
TRAINING_DIR = ROOT / "training"
|
| 31 |
TRAIN_SCRIPT = TRAINING_DIR / "train.py"
|
| 32 |
if not TRAIN_SCRIPT.exists():
|
requirements.txt
CHANGED
|
@@ -1,7 +1,7 @@
|
|
| 1 |
gradio>=6.5,<7
|
| 2 |
torch>=2.3
|
| 3 |
torchvision>=0.18
|
| 4 |
-
transformers>=4.50
|
| 5 |
huggingface_hub>=0.25
|
| 6 |
Pillow>=10.0
|
| 7 |
numpy>=1.26
|
|
|
|
| 1 |
gradio>=6.5,<7
|
| 2 |
torch>=2.3
|
| 3 |
torchvision>=0.18
|
| 4 |
+
transformers>=4.50,<5
|
| 5 |
huggingface_hub>=0.25
|
| 6 |
Pillow>=10.0
|
| 7 |
numpy>=1.26
|
training/train.py
CHANGED
|
@@ -86,6 +86,36 @@ def evaluate(model, loader, device):
|
|
| 86 |
model.train()
|
| 87 |
return total/max(n,1)
|
| 88 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 89 |
def main():
|
| 90 |
p=argparse.ArgumentParser()
|
| 91 |
p.add_argument("--train-dir",required=True); p.add_argument("--val-dir",required=True)
|
|
@@ -114,6 +144,7 @@ def main():
|
|
| 114 |
)
|
| 115 |
device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 116 |
model.to(device)
|
|
|
|
| 117 |
|
| 118 |
tr=DataLoader(train,batch_size=a.batch_size,shuffle=True,num_workers=0,collate_fn=collate_fn)
|
| 119 |
va=DataLoader(val,batch_size=a.batch_size,shuffle=False,num_workers=0,collate_fn=collate_fn)
|
|
@@ -130,7 +161,12 @@ def main():
|
|
| 130 |
# RT-DETR's loss matcher uses nested target tensors (boxes/classes).
|
| 131 |
# Move every tensor in labels to the same device as the model.
|
| 132 |
if "labels" in batch:
|
| 133 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 134 |
loss=model(**batch).loss
|
| 135 |
loss.backward(); opt.step(); opt.zero_grad(set_to_none=True)
|
| 136 |
running += float(loss.item())
|
|
|
|
| 86 |
model.train()
|
| 87 |
return total/max(n,1)
|
| 88 |
|
| 89 |
+
|
| 90 |
+
def patch_rtdetr_denoising_device():
|
| 91 |
+
"""Work around RT-DETR denoising code creating CPU index tensors on some Transformers releases."""
|
| 92 |
+
try:
|
| 93 |
+
import transformers.models.rt_detr.modeling_rt_detr as rtdetr_mod
|
| 94 |
+
original = rtdetr_mod.get_contrastive_denoising_training_group
|
| 95 |
+
if getattr(original, "_icecream_device_patch", False):
|
| 96 |
+
return
|
| 97 |
+
|
| 98 |
+
def wrapped(targets, *args, **kwargs):
|
| 99 |
+
# class_embed is the 4th positional argument in the supported RT-DETR versions.
|
| 100 |
+
class_embed = args[2] if len(args) >= 3 else kwargs.get("class_embed")
|
| 101 |
+
try:
|
| 102 |
+
device = next(class_embed.parameters()).device
|
| 103 |
+
except Exception:
|
| 104 |
+
device = None
|
| 105 |
+
if device is not None:
|
| 106 |
+
for target in targets:
|
| 107 |
+
if isinstance(target, dict):
|
| 108 |
+
for key in ("class_labels", "boxes"):
|
| 109 |
+
value = target.get(key)
|
| 110 |
+
if torch.is_tensor(value) and value.device != device:
|
| 111 |
+
target[key] = value.to(device)
|
| 112 |
+
return original(targets, *args, **kwargs)
|
| 113 |
+
|
| 114 |
+
wrapped._icecream_device_patch = True
|
| 115 |
+
rtdetr_mod.get_contrastive_denoising_training_group = wrapped
|
| 116 |
+
except Exception as exc:
|
| 117 |
+
print(f"Warning: RT-DETR denoising device patch was not installed: {exc}")
|
| 118 |
+
|
| 119 |
def main():
|
| 120 |
p=argparse.ArgumentParser()
|
| 121 |
p.add_argument("--train-dir",required=True); p.add_argument("--val-dir",required=True)
|
|
|
|
| 144 |
)
|
| 145 |
device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 146 |
model.to(device)
|
| 147 |
+
patch_rtdetr_denoising_device()
|
| 148 |
|
| 149 |
tr=DataLoader(train,batch_size=a.batch_size,shuffle=True,num_workers=0,collate_fn=collate_fn)
|
| 150 |
va=DataLoader(val,batch_size=a.batch_size,shuffle=False,num_workers=0,collate_fn=collate_fn)
|
|
|
|
| 161 |
# RT-DETR's loss matcher uses nested target tensors (boxes/classes).
|
| 162 |
# Move every tensor in labels to the same device as the model.
|
| 163 |
if "labels" in batch:
|
| 164 |
+
# RT-DETR expects every nested target tensor on the same device as the model.
|
| 165 |
+
for target in batch["labels"]:
|
| 166 |
+
if isinstance(target, dict):
|
| 167 |
+
for key, value in list(target.items()):
|
| 168 |
+
if torch.is_tensor(value):
|
| 169 |
+
target[key] = value.to(device)
|
| 170 |
loss=model(**batch).loss
|
| 171 |
loss.backward(); opt.step(); opt.zero_grad(set_to_none=True)
|
| 172 |
running += float(loss.item())
|