ibsocr1 commited on
Commit
ebcb8ff
·
verified ·
1 Parent(s): fcd343d

Upload 5 files

Browse files
Files changed (2) hide show
  1. app.py +1 -1
  2. training/train.py +8 -0
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 = "aW1wb3J0IGFyZ3BhcnNlCmltcG9ydCBqc29uCmZyb20gcGF0aGxpYiBpbXBvcnQgUGF0aAoKaW1wb3J0IHRvcmNoCmZyb20gUElMIGltcG9ydCBJbWFnZQpmcm9tIHRvcmNoLnV0aWxzLmRhdGEgaW1wb3J0IERhdGFzZXQsIERhdGFMb2FkZXIKZnJvbSB0cWRtIGltcG9ydCB0cWRtCmZyb20gdHJhbnNmb3JtZXJzIGltcG9ydCBSVERldHJJbWFnZVByb2Nlc3NvciwgUlREZXRyRm9yT2JqZWN0RGV0ZWN0aW9uCgpCQVNFX01PREVMID0gIlBla2luZ1UvcnRkZXRyX3I1MHZkIgoKZGVmIGxvYWRfY2xhc3NlcyhwYXRoKToKICAgIHJldHVybiBbeC5zdHJpcCgpIGZvciB4IGluIFBhdGgocGF0aCkucmVhZF90ZXh0KCkuc3BsaXRsaW5lcygpIGlmIHguc3RyaXAoKV0KCmNsYXNzIENPQ09EZXRlY3Rpb25EYXRhc2V0KERhdGFzZXQpOgogICAgZGVmIF9faW5pdF9fKHNlbGYsIGltYWdlX2RpciwgYW5ub3RhdGlvbl9maWxlLCBwcm9jZXNzb3IpOgogICAgICAgIHNlbGYuaW1hZ2VfZGlyID0gUGF0aChpbWFnZV9kaXIpCiAgICAgICAgc2VsZi5wcm9jZXNzb3IgPSBwcm9jZXNzb3IKICAgICAgICBjb2NvID0ganNvbi5sb2FkcyhQYXRoKGFubm90YXRpb25fZmlsZSkucmVhZF90ZXh0KCkpCiAgICAgICAgc2VsZi5pbWFnZXMgPSB7eFsiaWQiXTogeCBmb3IgeCBpbiBjb2NvWyJpbWFnZXMiXX0KICAgICAgICBjYXRzID0gc29ydGVkKGNvY29bImNhdGVnb3JpZXMiXSwga2V5PWxhbWJkYSB4OiB4WyJpZCJdKQogICAgICAgIHNlbGYuY2F0ZWdvcnlfaWRfdG9fbGFiZWwgPSB7Y1siaWQiXTogaSBmb3IgaSxjIGluIGVudW1lcmF0ZShjYXRzKX0KICAgICAgICBhbm5zID0ge30KICAgICAgICBmb3IgYSBpbiBjb2NvWyJhbm5vdGF0aW9ucyJdOgogICAgICAgICAgICBpZiBub3QgYS5nZXQoImlzY3Jvd2QiLCAwKToKICAgICAgICAgICAgICAgIGFubnMuc2V0ZGVmYXVsdChhWyJpbWFnZV9pZCJdLCBbXSkuYXBwZW5kKGEpCiAgICAgICAgc2VsZi5yZWNvcmRzID0gW10KICAgICAgICBmb3IgaW1hZ2VfaWQsIGluZm8gaW4gc2VsZi5pbWFnZXMuaXRlbXMoKToKICAgICAgICAgICAgc2VsZi5yZWNvcmRzLmFwcGVuZCh7CiAgICAgICAgICAgICAgICAiaW1hZ2VfaWQiOiBpbWFnZV9pZCwgImZpbGVfbmFtZSI6IGluZm9bImZpbGVfbmFtZSJdLAogICAgICAgICAgICAgICAgIndpZHRoIjogaW5mb1sid2lkdGgiXSwgImhlaWdodCI6IGluZm9bImhlaWdodCJdLAogICAgICAgICAgICAgICAgImFubm90YXRpb25zIjogYW5ucy5nZXQoaW1hZ2VfaWQsIFtdKQogICAgICAgICAgICB9KQoKICAgIGRlZiBfX2xlbl9fKHNlbGYpOiByZXR1cm4gbGVuKHNlbGYucmVjb3JkcykKCiAgICBkZWYgX19nZXRpdGVtX18oc2VsZiwgaWR4KToKICAgICAgICByID0gc2VsZi5yZWNvcmRzW2lkeF0KICAgICAgICBpbWFnZSA9IEltYWdlLm9wZW4oc2VsZi5pbWFnZV9kaXIgLyByWyJmaWxlX25hbWUiXSkuY29udmVydCgiUkdCIikKICAgICAgICBhbm5zID0gW10KICAgICAgICBmb3IgYSBpbiByWyJhbm5vdGF0aW9ucyJdOgogICAgICAgICAgICB4LHksdyxoID0gYVsiYmJveCJdCiAgICAgICAgICAgIGlmIHcgPD0gMCBvciBoIDw9IDA6IGNvbnRpbnVlCiAgICAgICAgICAgIGFubnMuYXBwZW5kKHsKICAgICAgICAgICAgICAgICJpZCI6IGFbImlkIl0sICJpbWFnZV9pZCI6IGludChpZHgpLAogICAgICAgICAgICAgICAgImNhdGVnb3J5X2lkIjogc2VsZi5jYXRlZ29yeV9pZF90b19sYWJlbFthWyJjYXRlZ29yeV9pZCJdXSwKICAgICAgICAgICAgICAgICJiYm94IjogW3gseSx3LGhdLCAiYXJlYSI6IGZsb2F0KGEuZ2V0KCJhcmVhIix3KmgpKSwKICAgICAgICAgICAgICAgICJpc2Nyb3dkIjogMAogICAgICAgICAgICB9KQogICAgICAgIGVuY29kZWQgPSBzZWxmLnByb2Nlc3NvcigKICAgICAgICAgICAgaW1hZ2VzPWltYWdlLAogICAgICAgICAgICBhbm5vdGF0aW9ucz17ImltYWdlX2lkIjogaW50KGlkeCksICJhbm5vdGF0aW9ucyI6IGFubnN9LAogICAgICAgICAgICByZXR1cm5fdGVuc29ycz0icHQiCiAgICAgICAgKQogICAgICAgIGVuY29kZWRbInBpeGVsX3ZhbHVlcyJdID0gZW5jb2RlZFsicGl4ZWxfdmFsdWVzIl0uc3F1ZWV6ZSgwKQogICAgICAgIGlmICJwaXhlbF9tYXNrIiBpbiBlbmNvZGVkOgogICAgICAgICAgICBlbmNvZGVkWyJwaXhlbF9tYXNrIl0gPSBlbmNvZGVkWyJwaXhlbF9tYXNrIl0uc3F1ZWV6ZSgwKQogICAgICAgIGVuY29kZWRbImxhYmVscyJdID0gZW5jb2RlZFsibGFiZWxzIl1bMF0KICAgICAgICByZXR1cm4gZW5jb2RlZAoKZGVmIGNvbGxhdGVfZm4oYmF0Y2gpOgogICAgb3V0ID0geyJwaXhlbF92YWx1ZXMiOiB0b3JjaC5zdGFjayhbeFsicGl4ZWxfdmFsdWVzIl0gZm9yIHggaW4gYmF0Y2hdKSwKICAgICAgICAgICAibGFiZWxzIjogW3hbImxhYmVscyJdIGZvciB4IGluIGJhdGNoXX0KICAgIGlmICJwaXhlbF9tYXNrIiBpbiBiYXRjaFswXToKICAgICAgICBvdXRbInBpeGVsX21hc2siXSA9IHRvcmNoLnN0YWNrKFt4WyJwaXhlbF9tYXNrIl0gZm9yIHggaW4gYmF0Y2hdKQogICAgcmV0dXJuIG91dAoKZGVmIG1vdmVfdG9fZGV2aWNlKG9iaiwgZGV2aWNlKToKICAgIGlmIHRvcmNoLmlzX3RlbnNvcihvYmopOgogICAgICAgIHJldHVybiBvYmoudG8oZGV2aWNlKQogICAgaWYgaXNpbnN0YW5jZShvYmosIGRpY3QpOgogICAgICAgIHJldHVybiB7azogbW92ZV90b19kZXZpY2UodiwgZGV2aWNlKSBmb3IgaywgdiBpbiBvYmouaXRlbXMoKX0KICAgIGlmIGlzaW5zdGFuY2Uob2JqLCBsaXN0KToKICAgICAgICByZXR1cm4gW21vdmVfdG9fZGV2aWNlKHYsIGRldmljZSkgZm9yIHYgaW4gb2JqXQogICAgaWYgaXNpbnN0YW5jZShvYmosIHR1cGxlKToKICAgICAgICByZXR1cm4gdHVwbGUobW92ZV90b19kZXZpY2UodiwgZGV2aWNlKSBmb3IgdiBpbiBvYmopCiAgICByZXR1cm4gb2JqCgpkZWYgZXZhbHVhdGUobW9kZWwsIGxvYWRlciwgZGV2aWNlKToKICAgIG1vZGVsLmV2YWwoKTsgdG90YWw9MDsgbj0wCiAgICB3aXRoIHRvcmNoLm5vX2dyYWQoKToKICAgICAgICBmb3IgYmF0Y2ggaW4gbG9hZGVyOgogICAgICAgICAgICBiYXRjaD1tb3ZlX3RvX2RldmljZShiYXRjaCwgZGV2aWNlKQogICAgICAgICAgICB0b3RhbCArPSBmbG9hdChtb2RlbCgqKmJhdGNoKS5sb3NzLml0ZW0oKSk7IG4gKz0gMQogICAgbW9kZWwudHJhaW4oKQogICAgcmV0dXJuIHRvdGFsL21heChuLDEpCgpkZWYgbWFpbigpOgogICAgcD1hcmdwYXJzZS5Bcmd1bWVudFBhcnNlcigpCiAgICBwLmFkZF9hcmd1bWVudCgiLS10cmFpbi1kaXIiLHJlcXVpcmVkPVRydWUpOyBwLmFkZF9hcmd1bWVudCgiLS12YWwtZGlyIixyZXF1aXJlZD1UcnVlKQogICAgcC5hZGRfYXJndW1lbnQoIi0tY2xhc3NlcyIscmVxdWlyZWQ9VHJ1ZSk7IHAuYWRkX2FyZ3VtZW50KCItLW91dHB1dC1kaXIiLGRlZmF1bHQ9Im1vZGVsIikKICAgIHAuYWRkX2FyZ3VtZW50KCItLWVwb2NocyIsdHlwZT1pbnQsZGVmYXVsdD0zMCk7IHAuYWRkX2FyZ3VtZW50KCItLWJhdGNoLXNpemUiLHR5cGU9aW50LGRlZmF1bHQ9MikKICAgIHAuYWRkX2FyZ3VtZW50KCItLWxlYXJuaW5nLXJhdGUiLHR5cGU9ZmxvYXQsZGVmYXVsdD0xZS01KTsgcC5hZGRfYXJndW1lbnQoIi0td2VpZ2h0LWRlY2F5Iix0eXBlPWZsb2F0LGRlZmF1bHQ9MWUtNCkKICAgIHAuYWRkX2FyZ3VtZW50KCItLW51bS13b3JrZXJzIix0eXBlPWludCxkZWZhdWx0PTIpCiAgICBhPXAucGFyc2VfYXJncygpCgogICAgY2xhc3Nlcz1sb2FkX2NsYXNzZXMoYS5jbGFzc2VzKQogICAgaWQybGFiZWw9e2k6biBmb3IgaSxuIGluIGVudW1lcmF0ZShjbGFzc2VzKX0KICAgIGxhYmVsMmlkPXtuOmkgZm9yIGksbiBpbiBlbnVtZXJhdGUoY2xhc3Nlcyl9CgogICAgcHJvYz1SVERldHJJbWFnZVByb2Nlc3Nvci5mcm9tX3ByZXRyYWluZWQoQkFTRV9NT0RFTCkKICAgIHRyYWluPUNPQ09EZXRlY3Rpb25EYXRhc2V0KFBhdGgoYS50cmFpbl9kaXIpLyJpbWFnZXMiLFBhdGgoYS50cmFpbl9kaXIpLyJhbm5vdGF0aW9ucy5qc29uIixwcm9jKQogICAgdmFsPUNPQ09EZXRlY3Rpb25EYXRhc2V0KFBhdGgoYS52YWxfZGlyKS8iaW1hZ2VzIixQYXRoKGEudmFsX2RpcikvImFubm90YXRpb25zLmpzb24iLHByb2MpCgogICAgaWYgbGVuKHRyYWluKT09MCBvciBsZW4odmFsKT09MDoKICAgICAgICByYWlzZSBWYWx1ZUVycm9yKCJUcmFpbmluZyBhbmQgdmFsaWRhdGlvbiBkYXRhc2V0cyBtdXN0IGNvbnRhaW4gYXQgbGVhc3Qgb25lIGltYWdlLiIpCiAgICBpZiBsZW4odHJhaW4uY2F0ZWdvcnlfaWRfdG9fbGFiZWwpIT1sZW4oY2xhc3Nlcykgb3IgbGVuKHZhbC5jYXRlZ29yeV9pZF90b19sYWJlbCkhPWxlbihjbGFzc2VzKToKICAgICAgICByYWlzZSBWYWx1ZUVycm9yKCJDT0NPIGNhdGVnb3JpZXMgZG8gbm90IG1hdGNoIGNsYXNzZXMudHh0LiBSZWJ1aWxkIHRoZSBkYXRhc2V0IGFmdGVyIHNhdmluZyB0aGUgY2xhc3Nlcy4iKQoKICAgIG1vZGVsPVJURGV0ckZvck9iamVjdERldGVjdGlvbi5mcm9tX3ByZXRyYWluZWQoCiAgICAgICAgQkFTRV9NT0RFTCxudW1fbGFiZWxzPWxlbihjbGFzc2VzKSxpZDJsYWJlbD1pZDJsYWJlbCxsYWJlbDJpZD1sYWJlbDJpZCwKICAgICAgICBpZ25vcmVfbWlzbWF0Y2hlZF9zaXplcz1UcnVlCiAgICApCiAgICBkZXZpY2U9dG9yY2guZGV2aWNlKCJjdWRhIiBpZiB0b3JjaC5jdWRhLmlzX2F2YWlsYWJsZSgpIGVsc2UgImNwdSIpCiAgICBtb2RlbC50byhkZXZpY2UpCgogICAgdHI9RGF0YUxvYWRlcih0cmFpbixiYXRjaF9zaXplPWEuYmF0Y2hfc2l6ZSxzaHVmZmxlPVRydWUsbnVtX3dvcmtlcnM9MCxjb2xsYXRlX2ZuPWNvbGxhdGVfZm4pCiAgICB2YT1EYXRhTG9hZGVyKHZhbCxiYXRjaF9zaXplPWEuYmF0Y2hfc2l6ZSxzaHVmZmxlPUZhbHNlLG51bV93b3JrZXJzPTAsY29sbGF0ZV9mbj1jb2xsYXRlX2ZuKQogICAgb3B0PXRvcmNoLm9wdGltLkFkYW1XKG1vZGVsLnBhcmFtZXRlcnMoKSxscj1hLmxlYXJuaW5nX3JhdGUsd2VpZ2h0X2RlY2F5PWEud2VpZ2h0X2RlY2F5KQoKICAgIG91dGRpcj1QYXRoKGEub3V0cHV0X2Rpcik7IG91dGRpci5ta2RpcihwYXJlbnRzPVRydWUsZXhpc3Rfb2s9VHJ1ZSkKICAgIGJlc3Q9ZmxvYXQoImluZiIpCgogICAgZm9yIGVwb2NoIGluIHJhbmdlKGEuZXBvY2hzKToKICAgICAgICBtb2RlbC50cmFpbigpOyBydW5uaW5nPTAKICAgICAgICBiYXI9dHFkbSh0cixkZXNjPWYiZXBvY2gge2Vwb2NoKzF9L3thLmVwb2Noc30iKQogICAgICAgIGZvciBzdGVwLGJhdGNoIGluIGVudW1lcmF0ZShiYXIpOgogICAgICAgICAgICBiYXRjaD1tb3ZlX3RvX2RldmljZShiYXRjaCwgZGV2aWNlKQogICAgICAgICAgICBsb3NzPW1vZGVsKCoqYmF0Y2gpLmxvc3MKICAgICAgICAgICAgbG9zcy5iYWNrd2FyZCgpOyBvcHQuc3RlcCgpOyBvcHQuemVyb19ncmFkKHNldF90b19ub25lPVRydWUpCiAgICAgICAgICAgIHJ1bm5pbmcgKz0gZmxvYXQobG9zcy5pdGVtKCkpCiAgICAgICAgICAgIGJhci5zZXRfcG9zdGZpeChsb3NzPWYie3J1bm5pbmcvKHN0ZXArMSk6LjRmfSIpCiAgICAgICAgdmw9ZXZhbHVhdGUobW9kZWwsdmEsZGV2aWNlKQogICAgICAgIHByaW50KGYidmFsaWRhdGlvbl9sb3NzPXt2bDouNGZ9IikKICAgICAgICBpZiB2bDxiZXN0OgogICAgICAgICAgICBiZXN0PXZsCiAgICAgICAgICAgIG1vZGVsLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpCiAgICAgICAgICAgIHByb2Muc2F2ZV9wcmV0cmFpbmVkKG91dGRpcikKICAgICAgICAgICAgKG91dGRpci8iY2xhc3Nlcy5qc29uIikud3JpdGVfdGV4dChqc29uLmR1bXBzKHsiaWQybGFiZWwiOmlkMmxhYmVsLCJsYWJlbDJpZCI6bGFiZWwyaWR9LGluZGVudD0yKSkKICAgIG1vZGVsLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpOyBwcm9jLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpCgppZiBfX25hbWVfXz09Il9fbWFpbl9fIjogbWFpbigpCg=="
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 = "aW1wb3J0IGFyZ3BhcnNlCmltcG9ydCBqc29uCmZyb20gcGF0aGxpYiBpbXBvcnQgUGF0aAoKaW1wb3J0IHRvcmNoCmZyb20gUElMIGltcG9ydCBJbWFnZQpmcm9tIHRvcmNoLnV0aWxzLmRhdGEgaW1wb3J0IERhdGFzZXQsIERhdGFMb2FkZXIKZnJvbSB0cWRtIGltcG9ydCB0cWRtCmZyb20gdHJhbnNmb3JtZXJzIGltcG9ydCBSVERldHJJbWFnZVByb2Nlc3NvciwgUlREZXRyRm9yT2JqZWN0RGV0ZWN0aW9uCgpCQVNFX01PREVMID0gIlBla2luZ1UvcnRkZXRyX3I1MHZkIgoKZGVmIGxvYWRfY2xhc3NlcyhwYXRoKToKICAgIHJldHVybiBbeC5zdHJpcCgpIGZvciB4IGluIFBhdGgocGF0aCkucmVhZF90ZXh0KCkuc3BsaXRsaW5lcygpIGlmIHguc3RyaXAoKV0KCmNsYXNzIENPQ09EZXRlY3Rpb25EYXRhc2V0KERhdGFzZXQpOgogICAgZGVmIF9faW5pdF9fKHNlbGYsIGltYWdlX2RpciwgYW5ub3RhdGlvbl9maWxlLCBwcm9jZXNzb3IpOgogICAgICAgIHNlbGYuaW1hZ2VfZGlyID0gUGF0aChpbWFnZV9kaXIpCiAgICAgICAgc2VsZi5wcm9jZXNzb3IgPSBwcm9jZXNzb3IKICAgICAgICBjb2NvID0ganNvbi5sb2FkcyhQYXRoKGFubm90YXRpb25fZmlsZSkucmVhZF90ZXh0KCkpCiAgICAgICAgc2VsZi5pbWFnZXMgPSB7eFsiaWQiXTogeCBmb3IgeCBpbiBjb2NvWyJpbWFnZXMiXX0KICAgICAgICBjYXRzID0gc29ydGVkKGNvY29bImNhdGVnb3JpZXMiXSwga2V5PWxhbWJkYSB4OiB4WyJpZCJdKQogICAgICAgIHNlbGYuY2F0ZWdvcnlfaWRfdG9fbGFiZWwgPSB7Y1siaWQiXTogaSBmb3IgaSxjIGluIGVudW1lcmF0ZShjYXRzKX0KICAgICAgICBhbm5zID0ge30KICAgICAgICBmb3IgYSBpbiBjb2NvWyJhbm5vdGF0aW9ucyJdOgogICAgICAgICAgICBpZiBub3QgYS5nZXQoImlzY3Jvd2QiLCAwKToKICAgICAgICAgICAgICAgIGFubnMuc2V0ZGVmYXVsdChhWyJpbWFnZV9pZCJdLCBbXSkuYXBwZW5kKGEpCiAgICAgICAgc2VsZi5yZWNvcmRzID0gW10KICAgICAgICBmb3IgaW1hZ2VfaWQsIGluZm8gaW4gc2VsZi5pbWFnZXMuaXRlbXMoKToKICAgICAgICAgICAgc2VsZi5yZWNvcmRzLmFwcGVuZCh7CiAgICAgICAgICAgICAgICAiaW1hZ2VfaWQiOiBpbWFnZV9pZCwgImZpbGVfbmFtZSI6IGluZm9bImZpbGVfbmFtZSJdLAogICAgICAgICAgICAgICAgIndpZHRoIjogaW5mb1sid2lkdGgiXSwgImhlaWdodCI6IGluZm9bImhlaWdodCJdLAogICAgICAgICAgICAgICAgImFubm90YXRpb25zIjogYW5ucy5nZXQoaW1hZ2VfaWQsIFtdKQogICAgICAgICAgICB9KQoKICAgIGRlZiBfX2xlbl9fKHNlbGYpOiByZXR1cm4gbGVuKHNlbGYucmVjb3JkcykKCiAgICBkZWYgX19nZXRpdGVtX18oc2VsZiwgaWR4KToKICAgICAgICByID0gc2VsZi5yZWNvcmRzW2lkeF0KICAgICAgICBpbWFnZSA9IEltYWdlLm9wZW4oc2VsZi5pbWFnZV9kaXIgLyByWyJmaWxlX25hbWUiXSkuY29udmVydCgiUkdCIikKICAgICAgICBhbm5zID0gW10KICAgICAgICBmb3IgYSBpbiByWyJhbm5vdGF0aW9ucyJdOgogICAgICAgICAgICB4LHksdyxoID0gYVsiYmJveCJdCiAgICAgICAgICAgIGlmIHcgPD0gMCBvciBoIDw9IDA6IGNvbnRpbnVlCiAgICAgICAgICAgIGFubnMuYXBwZW5kKHsKICAgICAgICAgICAgICAgICJpZCI6IGFbImlkIl0sICJpbWFnZV9pZCI6IGludChpZHgpLAogICAgICAgICAgICAgICAgImNhdGVnb3J5X2lkIjogc2VsZi5jYXRlZ29yeV9pZF90b19sYWJlbFthWyJjYXRlZ29yeV9pZCJdXSwKICAgICAgICAgICAgICAgICJiYm94IjogW3gseSx3LGhdLCAiYXJlYSI6IGZsb2F0KGEuZ2V0KCJhcmVhIix3KmgpKSwKICAgICAgICAgICAgICAgICJpc2Nyb3dkIjogMAogICAgICAgICAgICB9KQogICAgICAgIGVuY29kZWQgPSBzZWxmLnByb2Nlc3NvcigKICAgICAgICAgICAgaW1hZ2VzPWltYWdlLAogICAgICAgICAgICBhbm5vdGF0aW9ucz17ImltYWdlX2lkIjogaW50KGlkeCksICJhbm5vdGF0aW9ucyI6IGFubnN9LAogICAgICAgICAgICByZXR1cm5fdGVuc29ycz0icHQiCiAgICAgICAgKQogICAgICAgIGVuY29kZWRbInBpeGVsX3ZhbHVlcyJdID0gZW5jb2RlZFsicGl4ZWxfdmFsdWVzIl0uc3F1ZWV6ZSgwKQogICAgICAgIGlmICJwaXhlbF9tYXNrIiBpbiBlbmNvZGVkOgogICAgICAgICAgICBlbmNvZGVkWyJwaXhlbF9tYXNrIl0gPSBlbmNvZGVkWyJwaXhlbF9tYXNrIl0uc3F1ZWV6ZSgwKQogICAgICAgIGVuY29kZWRbImxhYmVscyJdID0gZW5jb2RlZFsibGFiZWxzIl1bMF0KICAgICAgICByZXR1cm4gZW5jb2RlZAoKZGVmIGNvbGxhdGVfZm4oYmF0Y2gpOgogICAgb3V0ID0geyJwaXhlbF92YWx1ZXMiOiB0b3JjaC5zdGFjayhbeFsicGl4ZWxfdmFsdWVzIl0gZm9yIHggaW4gYmF0Y2hdKSwKICAgICAgICAgICAibGFiZWxzIjogW3hbImxhYmVscyJdIGZvciB4IGluIGJhdGNoXX0KICAgIGlmICJwaXhlbF9tYXNrIiBpbiBiYXRjaFswXToKICAgICAgICBvdXRbInBpeGVsX21hc2siXSA9IHRvcmNoLnN0YWNrKFt4WyJwaXhlbF9tYXNrIl0gZm9yIHggaW4gYmF0Y2hdKQogICAgcmV0dXJuIG91dAoKZGVmIG1vdmVfdG9fZGV2aWNlKG9iaiwgZGV2aWNlKToKICAgIGlmIHRvcmNoLmlzX3RlbnNvcihvYmopOgogICAgICAgIHJldHVybiBvYmoudG8oZGV2aWNlKQogICAgaWYgaXNpbnN0YW5jZShvYmosIGRpY3QpOgogICAgICAgIHJldHVybiB7azogbW92ZV90b19kZXZpY2UodiwgZGV2aWNlKSBmb3IgaywgdiBpbiBvYmouaXRlbXMoKX0KICAgIGlmIGlzaW5zdGFuY2Uob2JqLCBsaXN0KToKICAgICAgICByZXR1cm4gW21vdmVfdG9fZGV2aWNlKHYsIGRldmljZSkgZm9yIHYgaW4gb2JqXQogICAgaWYgaXNpbnN0YW5jZShvYmosIHR1cGxlKToKICAgICAgICByZXR1cm4gdHVwbGUobW92ZV90b19kZXZpY2UodiwgZGV2aWNlKSBmb3IgdiBpbiBvYmopCiAgICByZXR1cm4gb2JqCgpkZWYgZXZhbHVhdGUobW9kZWwsIGxvYWRlciwgZGV2aWNlKToKICAgIG1vZGVsLmV2YWwoKTsgdG90YWw9MDsgbj0wCiAgICB3aXRoIHRvcmNoLm5vX2dyYWQoKToKICAgICAgICBmb3IgYmF0Y2ggaW4gbG9hZGVyOgogICAgICAgICAgICBiYXRjaD1tb3ZlX3RvX2RldmljZShiYXRjaCwgZGV2aWNlKQogICAgICAgICAgICB0b3RhbCArPSBmbG9hdChtb2RlbCgqKmJhdGNoKS5sb3NzLml0ZW0oKSk7IG4gKz0gMQogICAgbW9kZWwudHJhaW4oKQogICAgcmV0dXJuIHRvdGFsL21heChuLDEpCgpkZWYgbWFpbigpOgogICAgcD1hcmdwYXJzZS5Bcmd1bWVudFBhcnNlcigpCiAgICBwLmFkZF9hcmd1bWVudCgiLS10cmFpbi1kaXIiLHJlcXVpcmVkPVRydWUpOyBwLmFkZF9hcmd1bWVudCgiLS12YWwtZGlyIixyZXF1aXJlZD1UcnVlKQogICAgcC5hZGRfYXJndW1lbnQoIi0tY2xhc3NlcyIscmVxdWlyZWQ9VHJ1ZSk7IHAuYWRkX2FyZ3VtZW50KCItLW91dHB1dC1kaXIiLGRlZmF1bHQ9Im1vZGVsIikKICAgIHAuYWRkX2FyZ3VtZW50KCItLWVwb2NocyIsdHlwZT1pbnQsZGVmYXVsdD0zMCk7IHAuYWRkX2FyZ3VtZW50KCItLWJhdGNoLXNpemUiLHR5cGU9aW50LGRlZmF1bHQ9MikKICAgIHAuYWRkX2FyZ3VtZW50KCItLWxlYXJuaW5nLXJhdGUiLHR5cGU9ZmxvYXQsZGVmYXVsdD0xZS01KTsgcC5hZGRfYXJndW1lbnQoIi0td2VpZ2h0LWRlY2F5Iix0eXBlPWZsb2F0LGRlZmF1bHQ9MWUtNCkKICAgIHAuYWRkX2FyZ3VtZW50KCItLW51bS13b3JrZXJzIix0eXBlPWludCxkZWZhdWx0PTIpCiAgICBhPXAucGFyc2VfYXJncygpCgogICAgY2xhc3Nlcz1sb2FkX2NsYXNzZXMoYS5jbGFzc2VzKQogICAgaWQybGFiZWw9e2k6biBmb3IgaSxuIGluIGVudW1lcmF0ZShjbGFzc2VzKX0KICAgIGxhYmVsMmlkPXtuOmkgZm9yIGksbiBpbiBlbnVtZXJhdGUoY2xhc3Nlcyl9CgogICAgcHJvYz1SVERldHJJbWFnZVByb2Nlc3Nvci5mcm9tX3ByZXRyYWluZWQoQkFTRV9NT0RFTCkKICAgIHRyYWluPUNPQ09EZXRlY3Rpb25EYXRhc2V0KFBhdGgoYS50cmFpbl9kaXIpLyJpbWFnZXMiLFBhdGgoYS50cmFpbl9kaXIpLyJhbm5vdGF0aW9ucy5qc29uIixwcm9jKQogICAgdmFsPUNPQ09EZXRlY3Rpb25EYXRhc2V0KFBhdGgoYS52YWxfZGlyKS8iaW1hZ2VzIixQYXRoKGEudmFsX2RpcikvImFubm90YXRpb25zLmpzb24iLHByb2MpCgogICAgaWYgbGVuKHRyYWluKT09MCBvciBsZW4odmFsKT09MDoKICAgICAgICByYWlzZSBWYWx1ZUVycm9yKCJUcmFpbmluZyBhbmQgdmFsaWRhdGlvbiBkYXRhc2V0cyBtdXN0IGNvbnRhaW4gYXQgbGVhc3Qgb25lIGltYWdlLiIpCiAgICBpZiBsZW4odHJhaW4uY2F0ZWdvcnlfaWRfdG9fbGFiZWwpIT1sZW4oY2xhc3Nlcykgb3IgbGVuKHZhbC5jYXRlZ29yeV9pZF90b19sYWJlbCkhPWxlbihjbGFzc2VzKToKICAgICAgICByYWlzZSBWYWx1ZUVycm9yKCJDT0NPIGNhdGVnb3JpZXMgZG8gbm90IG1hdGNoIGNsYXNzZXMudHh0LiBSZWJ1aWxkIHRoZSBkYXRhc2V0IGFmdGVyIHNhdmluZyB0aGUgY2xhc3Nlcy4iKQoKICAgIG1vZGVsPVJURGV0ckZvck9iamVjdERldGVjdGlvbi5mcm9tX3ByZXRyYWluZWQoCiAgICAgICAgQkFTRV9NT0RFTCxudW1fbGFiZWxzPWxlbihjbGFzc2VzKSxpZDJsYWJlbD1pZDJsYWJlbCxsYWJlbDJpZD1sYWJlbDJpZCwKICAgICAgICBpZ25vcmVfbWlzbWF0Y2hlZF9zaXplcz1UcnVlCiAgICApCiAgICAjIFRoZSBIdWdnaW5nIEZhY2UgUlQtREVUUiBpbXBsZW1lbnRhdGlvbiBjYW4gbGVhdmUgdGhlIGNvbnRyYXN0aXZlCiAgICAjIGRlbm9pc2luZyBjbGFzcy1pbmRleCB0ZW5zb3Igb24gQ1BVIHdoaWxlIHRoZSBtb2RlbCBpcyBvbiBDVURBLiBUaGlzCiAgICAjIGNhdXNlcyB0aGUgZW1iZWRkaW5nIGxvb2t1cCB0byBmYWlsIGR1cmluZyBmaW5lLXR1bmluZy4gRGlzYWJsZSB0aGUKICAgICMgb3B0aW9uYWwgZGVub2lzaW5nIGJyYW5jaDsgb3JkaW5hcnkgb2JqZWN0LWRldGVjdGlvbiB0cmFpbmluZyByZW1haW5zCiAgICAjIGVuYWJsZWQgYW5kIHdvcmtzIHdpdGggdGhlIDEwIGN1c3RvbSBjbGFzc2VzLgogICAgbW9kZWwuY29uZmlnLm51bV9kZW5vaXNpbmcgPSAwCiAgICBpZiBoYXNhdHRyKG1vZGVsLCAibW9kZWwiKSBhbmQgaGFzYXR0cihtb2RlbC5tb2RlbCwgImRlY29kZXIiKToKICAgICAgICBtb2RlbC5tb2RlbC5kZWNvZGVyLm51bV9kZW5vaXNpbmcgPSAwCiAgICBkZXZpY2U9dG9yY2guZGV2aWNlKCJjdWRhIiBpZiB0b3JjaC5jdWRhLmlzX2F2YWlsYWJsZSgpIGVsc2UgImNwdSIpCiAgICBtb2RlbC50byhkZXZpY2UpCgogICAgdHI9RGF0YUxvYWRlcih0cmFpbixiYXRjaF9zaXplPWEuYmF0Y2hfc2l6ZSxzaHVmZmxlPVRydWUsbnVtX3dvcmtlcnM9MCxjb2xsYXRlX2ZuPWNvbGxhdGVfZm4pCiAgICB2YT1EYXRhTG9hZGVyKHZhbCxiYXRjaF9zaXplPWEuYmF0Y2hfc2l6ZSxzaHVmZmxlPUZhbHNlLG51bV93b3JrZXJzPTAsY29sbGF0ZV9mbj1jb2xsYXRlX2ZuKQogICAgb3B0PXRvcmNoLm9wdGltLkFkYW1XKG1vZGVsLnBhcmFtZXRlcnMoKSxscj1hLmxlYXJuaW5nX3JhdGUsd2VpZ2h0X2RlY2F5PWEud2VpZ2h0X2RlY2F5KQoKICAgIG91dGRpcj1QYXRoKGEub3V0cHV0X2Rpcik7IG91dGRpci5ta2RpcihwYXJlbnRzPVRydWUsZXhpc3Rfb2s9VHJ1ZSkKICAgIGJlc3Q9ZmxvYXQoImluZiIpCgogICAgZm9yIGVwb2NoIGluIHJhbmdlKGEuZXBvY2hzKToKICAgICAgICBtb2RlbC50cmFpbigpOyBydW5uaW5nPTAKICAgICAgICBiYXI9dHFkbSh0cixkZXNjPWYiZXBvY2gge2Vwb2NoKzF9L3thLmVwb2Noc30iKQogICAgICAgIGZvciBzdGVwLGJhdGNoIGluIGVudW1lcmF0ZShiYXIpOgogICAgICAgICAgICBiYXRjaD1tb3ZlX3RvX2RldmljZShiYXRjaCwgZGV2aWNlKQogICAgICAgICAgICBsb3NzPW1vZGVsKCoqYmF0Y2gpLmxvc3MKICAgICAgICAgICAgbG9zcy5iYWNrd2FyZCgpOyBvcHQuc3RlcCgpOyBvcHQuemVyb19ncmFkKHNldF90b19ub25lPVRydWUpCiAgICAgICAgICAgIHJ1bm5pbmcgKz0gZmxvYXQobG9zcy5pdGVtKCkpCiAgICAgICAgICAgIGJhci5zZXRfcG9zdGZpeChsb3NzPWYie3J1bm5pbmcvKHN0ZXArMSk6LjRmfSIpCiAgICAgICAgdmw9ZXZhbHVhdGUobW9kZWwsdmEsZGV2aWNlKQogICAgICAgIHByaW50KGYidmFsaWRhdGlvbl9sb3NzPXt2bDouNGZ9IikKICAgICAgICBpZiB2bDxiZXN0OgogICAgICAgICAgICBiZXN0PXZsCiAgICAgICAgICAgIG1vZGVsLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpCiAgICAgICAgICAgIHByb2Muc2F2ZV9wcmV0cmFpbmVkKG91dGRpcikKICAgICAgICAgICAgKG91dGRpci8iY2xhc3Nlcy5qc29uIikud3JpdGVfdGV4dChqc29uLmR1bXBzKHsiaWQybGFiZWwiOmlkMmxhYmVsLCJsYWJlbDJpZCI6bGFiZWwyaWR9LGluZGVudD0yKSkKICAgIG1vZGVsLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpOyBwcm9jLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpCgppZiBfX25hbWVfXz09Il9fbWFpbl9fIjogbWFpbigpCg=="
30
  TRAINING_DIR = ROOT / "training"
31
  TRAIN_SCRIPT = TRAINING_DIR / "train.py"
32
  if not TRAIN_SCRIPT.exists():
training/train.py CHANGED
@@ -112,6 +112,14 @@ def main():
112
  BASE_MODEL,num_labels=len(classes),id2label=id2label,label2id=label2id,
113
  ignore_mismatched_sizes=True
114
  )
 
 
 
 
 
 
 
 
115
  device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
116
  model.to(device)
117
 
 
112
  BASE_MODEL,num_labels=len(classes),id2label=id2label,label2id=label2id,
113
  ignore_mismatched_sizes=True
114
  )
115
+ # The Hugging Face RT-DETR implementation can leave the contrastive
116
+ # denoising class-index tensor on CPU while the model is on CUDA. This
117
+ # causes the embedding lookup to fail during fine-tuning. Disable the
118
+ # optional denoising branch; ordinary object-detection training remains
119
+ # enabled and works with the 10 custom classes.
120
+ model.config.num_denoising = 0
121
+ if hasattr(model, "model") and hasattr(model.model, "decoder"):
122
+ model.model.decoder.num_denoising = 0
123
  device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
124
  model.to(device)
125