ibsocr1 commited on
Commit
bcf1d8f
·
verified ·
1 Parent(s): 123ef0f

Upload 5 files

Browse files
Files changed (4) hide show
  1. README.md +14 -0
  2. app.py +1 -1
  3. requirements.txt +1 -1
  4. 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
- batch["labels"] = move_to_device(batch["labels"], device)
 
 
 
 
 
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())