Spaces:
Running on Zero
Running on Zero
| import threading | |
| import io | |
| import json | |
| import os | |
| import shutil | |
| import subprocess | |
| import sys | |
| from collections import Counter | |
| from pathlib import Path | |
| import spaces | |
| import gradio as gr | |
| import torch | |
| from PIL import Image, ImageDraw | |
| from transformers import RTDetrForObjectDetection, RTDetrImageProcessor | |
| # Hugging Face Spaces can mount persistent storage at /data. | |
| # DATA_DIR can be overridden in Space Settings -> Variables. | |
| if os.getenv("DATA_DIR"): | |
| BASE = Path(os.environ["DATA_DIR"]) | |
| elif Path("/data").exists() and os.access("/data", os.W_OK): | |
| BASE = Path("/data") / "icecream_counter" | |
| else: | |
| BASE = Path("./data") | |
| ROOT = Path(__file__).resolve().parent | |
| # Self-heal the training script if a deployment omitted the training/ directory. | |
| _TRAINING_SCRIPT_B64 = "aW1wb3J0IGFyZ3BhcnNlCmltcG9ydCBqc29uCmZyb20gcGF0aGxpYiBpbXBvcnQgUGF0aAoKaW1wb3J0IHRvcmNoCmZyb20gUElMIGltcG9ydCBJbWFnZQpmcm9tIHRvcmNoLnV0aWxzLmRhdGEgaW1wb3J0IERhdGFzZXQsIERhdGFMb2FkZXIKZnJvbSB0cWRtIGltcG9ydCB0cWRtCmZyb20gdHJhbnNmb3JtZXJzIGltcG9ydCBSVERldHJJbWFnZVByb2Nlc3NvciwgUlREZXRyRm9yT2JqZWN0RGV0ZWN0aW9uCgpCQVNFX01PREVMID0gIlBla2luZ1UvcnRkZXRyX3I1MHZkIgoKZGVmIGxvYWRfY2xhc3NlcyhwYXRoKToKICAgIHJldHVybiBbeC5zdHJpcCgpIGZvciB4IGluIFBhdGgocGF0aCkucmVhZF90ZXh0KCkuc3BsaXRsaW5lcygpIGlmIHguc3RyaXAoKV0KCmNsYXNzIENPQ09EZXRlY3Rpb25EYXRhc2V0KERhdGFzZXQpOgogICAgZGVmIF9faW5pdF9fKHNlbGYsIGltYWdlX2RpciwgYW5ub3RhdGlvbl9maWxlLCBwcm9jZXNzb3IpOgogICAgICAgIHNlbGYuaW1hZ2VfZGlyID0gUGF0aChpbWFnZV9kaXIpCiAgICAgICAgc2VsZi5wcm9jZXNzb3IgPSBwcm9jZXNzb3IKICAgICAgICBjb2NvID0ganNvbi5sb2FkcyhQYXRoKGFubm90YXRpb25fZmlsZSkucmVhZF90ZXh0KCkpCiAgICAgICAgc2VsZi5pbWFnZXMgPSB7eFsiaWQiXTogeCBmb3IgeCBpbiBjb2NvWyJpbWFnZXMiXX0KICAgICAgICBjYXRzID0gc29ydGVkKGNvY29bImNhdGVnb3JpZXMiXSwga2V5PWxhbWJkYSB4OiB4WyJpZCJdKQogICAgICAgIHNlbGYuY2F0ZWdvcnlfaWRfdG9fbGFiZWwgPSB7Y1siaWQiXTogaSBmb3IgaSxjIGluIGVudW1lcmF0ZShjYXRzKX0KICAgICAgICBhbm5zID0ge30KICAgICAgICBmb3IgYSBpbiBjb2NvWyJhbm5vdGF0aW9ucyJdOgogICAgICAgICAgICBpZiBub3QgYS5nZXQoImlzY3Jvd2QiLCAwKToKICAgICAgICAgICAgICAgIGFubnMuc2V0ZGVmYXVsdChhWyJpbWFnZV9pZCJdLCBbXSkuYXBwZW5kKGEpCiAgICAgICAgc2VsZi5yZWNvcmRzID0gW10KICAgICAgICBmb3IgaW1hZ2VfaWQsIGluZm8gaW4gc2VsZi5pbWFnZXMuaXRlbXMoKToKICAgICAgICAgICAgc2VsZi5yZWNvcmRzLmFwcGVuZCh7CiAgICAgICAgICAgICAgICAiaW1hZ2VfaWQiOiBpbWFnZV9pZCwgImZpbGVfbmFtZSI6IGluZm9bImZpbGVfbmFtZSJdLAogICAgICAgICAgICAgICAgIndpZHRoIjogaW5mb1sid2lkdGgiXSwgImhlaWdodCI6IGluZm9bImhlaWdodCJdLAogICAgICAgICAgICAgICAgImFubm90YXRpb25zIjogYW5ucy5nZXQoaW1hZ2VfaWQsIFtdKQogICAgICAgICAgICB9KQoKICAgIGRlZiBfX2xlbl9fKHNlbGYpOiByZXR1cm4gbGVuKHNlbGYucmVjb3JkcykKCiAgICBkZWYgX19nZXRpdGVtX18oc2VsZiwgaWR4KToKICAgICAgICByID0gc2VsZi5yZWNvcmRzW2lkeF0KICAgICAgICBpbWFnZSA9IEltYWdlLm9wZW4oc2VsZi5pbWFnZV9kaXIgLyByWyJmaWxlX25hbWUiXSkuY29udmVydCgiUkdCIikKICAgICAgICBhbm5zID0gW10KICAgICAgICBmb3IgYSBpbiByWyJhbm5vdGF0aW9ucyJdOgogICAgICAgICAgICB4LHksdyxoID0gYVsiYmJveCJdCiAgICAgICAgICAgIGlmIHcgPD0gMCBvciBoIDw9IDA6IGNvbnRpbnVlCiAgICAgICAgICAgIGFubnMuYXBwZW5kKHsKICAgICAgICAgICAgICAgICJpZCI6IGFbImlkIl0sICJpbWFnZV9pZCI6IGludChpZHgpLAogICAgICAgICAgICAgICAgImNhdGVnb3J5X2lkIjogc2VsZi5jYXRlZ29yeV9pZF90b19sYWJlbFthWyJjYXRlZ29yeV9pZCJdXSwKICAgICAgICAgICAgICAgICJiYm94IjogW3gseSx3LGhdLCAiYXJlYSI6IGZsb2F0KGEuZ2V0KCJhcmVhIix3KmgpKSwKICAgICAgICAgICAgICAgICJpc2Nyb3dkIjogMAogICAgICAgICAgICB9KQogICAgICAgIGVuY29kZWQgPSBzZWxmLnByb2Nlc3NvcigKICAgICAgICAgICAgaW1hZ2VzPWltYWdlLAogICAgICAgICAgICBhbm5vdGF0aW9ucz17ImltYWdlX2lkIjogaW50KGlkeCksICJhbm5vdGF0aW9ucyI6IGFubnN9LAogICAgICAgICAgICByZXR1cm5fdGVuc29ycz0icHQiCiAgICAgICAgKQogICAgICAgIGVuY29kZWRbInBpeGVsX3ZhbHVlcyJdID0gZW5jb2RlZFsicGl4ZWxfdmFsdWVzIl0uc3F1ZWV6ZSgwKQogICAgICAgIGlmICJwaXhlbF9tYXNrIiBpbiBlbmNvZGVkOgogICAgICAgICAgICBlbmNvZGVkWyJwaXhlbF9tYXNrIl0gPSBlbmNvZGVkWyJwaXhlbF9tYXNrIl0uc3F1ZWV6ZSgwKQogICAgICAgIGVuY29kZWRbImxhYmVscyJdID0gZW5jb2RlZFsibGFiZWxzIl1bMF0KICAgICAgICByZXR1cm4gZW5jb2RlZAoKZGVmIGNvbGxhdGVfZm4oYmF0Y2gpOgogICAgb3V0ID0geyJwaXhlbF92YWx1ZXMiOiB0b3JjaC5zdGFjayhbeFsicGl4ZWxfdmFsdWVzIl0gZm9yIHggaW4gYmF0Y2hdKSwKICAgICAgICAgICAibGFiZWxzIjogW3hbImxhYmVscyJdIGZvciB4IGluIGJhdGNoXX0KICAgIGlmICJwaXhlbF9tYXNrIiBpbiBiYXRjaFswXToKICAgICAgICBvdXRbInBpeGVsX21hc2siXSA9IHRvcmNoLnN0YWNrKFt4WyJwaXhlbF9tYXNrIl0gZm9yIHggaW4gYmF0Y2hdKQogICAgcmV0dXJuIG91dAoKZGVmIG1vdmVfdG9fZGV2aWNlKG9iaiwgZGV2aWNlKToKICAgIGlmIHRvcmNoLmlzX3RlbnNvcihvYmopOgogICAgICAgIHJldHVybiBvYmoudG8oZGV2aWNlKQogICAgaWYgaXNpbnN0YW5jZShvYmosIGRpY3QpOgogICAgICAgIHJldHVybiB7azogbW92ZV90b19kZXZpY2UodiwgZGV2aWNlKSBmb3IgaywgdiBpbiBvYmouaXRlbXMoKX0KICAgIGlmIGlzaW5zdGFuY2Uob2JqLCBsaXN0KToKICAgICAgICByZXR1cm4gW21vdmVfdG9fZGV2aWNlKHYsIGRldmljZSkgZm9yIHYgaW4gb2JqXQogICAgaWYgaXNpbnN0YW5jZShvYmosIHR1cGxlKToKICAgICAgICByZXR1cm4gdHVwbGUobW92ZV90b19kZXZpY2UodiwgZGV2aWNlKSBmb3IgdiBpbiBvYmopCiAgICByZXR1cm4gb2JqCgpkZWYgZXZhbHVhdGUobW9kZWwsIGxvYWRlciwgZGV2aWNlKToKICAgIG1vZGVsLmV2YWwoKTsgdG90YWw9MDsgbj0wCiAgICB3aXRoIHRvcmNoLm5vX2dyYWQoKToKICAgICAgICBmb3IgYmF0Y2ggaW4gbG9hZGVyOgogICAgICAgICAgICBiYXRjaD1tb3ZlX3RvX2RldmljZShiYXRjaCwgZGV2aWNlKQogICAgICAgICAgICB0b3RhbCArPSBmbG9hdChtb2RlbCgqKmJhdGNoKS5sb3NzLml0ZW0oKSk7IG4gKz0gMQogICAgbW9kZWwudHJhaW4oKQogICAgcmV0dXJuIHRvdGFsL21heChuLDEpCgoKZGVmIHBhdGNoX3J0ZGV0cl9kZW5vaXNpbmdfZGV2aWNlKCk6CiAgICAiIiJXb3JrIGFyb3VuZCBSVC1ERVRSIGRlbm9pc2luZyBjb2RlIGNyZWF0aW5nIENQVSBpbmRleCB0ZW5zb3JzIG9uIHNvbWUgVHJhbnNmb3JtZXJzIHJlbGVhc2VzLiIiIgogICAgdHJ5OgogICAgICAgIGltcG9ydCB0cmFuc2Zvcm1lcnMubW9kZWxzLnJ0X2RldHIubW9kZWxpbmdfcnRfZGV0ciBhcyBydGRldHJfbW9kCiAgICAgICAgb3JpZ2luYWwgPSBydGRldHJfbW9kLmdldF9jb250cmFzdGl2ZV9kZW5vaXNpbmdfdHJhaW5pbmdfZ3JvdXAKICAgICAgICBpZiBnZXRhdHRyKG9yaWdpbmFsLCAiX2ljZWNyZWFtX2RldmljZV9wYXRjaCIsIEZhbHNlKToKICAgICAgICAgICAgcmV0dXJuCgogICAgICAgIGRlZiB3cmFwcGVkKHRhcmdldHMsICphcmdzLCAqKmt3YXJncyk6CiAgICAgICAgICAgICMgY2xhc3NfZW1iZWQgaXMgdGhlIDR0aCBwb3NpdGlvbmFsIGFyZ3VtZW50IGluIHRoZSBzdXBwb3J0ZWQgUlQtREVUUiB2ZXJzaW9ucy4KICAgICAgICAgICAgY2xhc3NfZW1iZWQgPSBhcmdzWzJdIGlmIGxlbihhcmdzKSA+PSAzIGVsc2Uga3dhcmdzLmdldCgiY2xhc3NfZW1iZWQiKQogICAgICAgICAgICB0cnk6CiAgICAgICAgICAgICAgICBkZXZpY2UgPSBuZXh0KGNsYXNzX2VtYmVkLnBhcmFtZXRlcnMoKSkuZGV2aWNlCiAgICAgICAgICAgIGV4Y2VwdCBFeGNlcHRpb246CiAgICAgICAgICAgICAgICBkZXZpY2UgPSBOb25lCiAgICAgICAgICAgIGlmIGRldmljZSBpcyBub3QgTm9uZToKICAgICAgICAgICAgICAgIGZvciB0YXJnZXQgaW4gdGFyZ2V0czoKICAgICAgICAgICAgICAgICAgICBpZiBpc2luc3RhbmNlKHRhcmdldCwgZGljdCk6CiAgICAgICAgICAgICAgICAgICAgICAgIGZvciBrZXkgaW4gKCJjbGFzc19sYWJlbHMiLCAiYm94ZXMiKToKICAgICAgICAgICAgICAgICAgICAgICAgICAgIHZhbHVlID0gdGFyZ2V0LmdldChrZXkpCiAgICAgICAgICAgICAgICAgICAgICAgICAgICBpZiB0b3JjaC5pc190ZW5zb3IodmFsdWUpIGFuZCB2YWx1ZS5kZXZpY2UgIT0gZGV2aWNlOgogICAgICAgICAgICAgICAgICAgICAgICAgICAgICAgIHRhcmdldFtrZXldID0gdmFsdWUudG8oZGV2aWNlKQogICAgICAgICAgICByZXR1cm4gb3JpZ2luYWwodGFyZ2V0cywgKmFyZ3MsICoqa3dhcmdzKQoKICAgICAgICB3cmFwcGVkLl9pY2VjcmVhbV9kZXZpY2VfcGF0Y2ggPSBUcnVlCiAgICAgICAgcnRkZXRyX21vZC5nZXRfY29udHJhc3RpdmVfZGVub2lzaW5nX3RyYWluaW5nX2dyb3VwID0gd3JhcHBlZAogICAgZXhjZXB0IEV4Y2VwdGlvbiBhcyBleGM6CiAgICAgICAgcHJpbnQoZiJXYXJuaW5nOiBSVC1ERVRSIGRlbm9pc2luZyBkZXZpY2UgcGF0Y2ggd2FzIG5vdCBpbnN0YWxsZWQ6IHtleGN9IikKCmRlZiBtYWluKCk6CiAgICBwPWFyZ3BhcnNlLkFyZ3VtZW50UGFyc2VyKCkKICAgIHAuYWRkX2FyZ3VtZW50KCItLXRyYWluLWRpciIscmVxdWlyZWQ9VHJ1ZSk7IHAuYWRkX2FyZ3VtZW50KCItLXZhbC1kaXIiLHJlcXVpcmVkPVRydWUpCiAgICBwLmFkZF9hcmd1bWVudCgiLS1jbGFzc2VzIixyZXF1aXJlZD1UcnVlKTsgcC5hZGRfYXJndW1lbnQoIi0tb3V0cHV0LWRpciIsZGVmYXVsdD0ibW9kZWwiKQogICAgcC5hZGRfYXJndW1lbnQoIi0tZXBvY2hzIix0eXBlPWludCxkZWZhdWx0PTMwKTsgcC5hZGRfYXJndW1lbnQoIi0tYmF0Y2gtc2l6ZSIsdHlwZT1pbnQsZGVmYXVsdD0yKQogICAgcC5hZGRfYXJndW1lbnQoIi0tbGVhcm5pbmctcmF0ZSIsdHlwZT1mbG9hdCxkZWZhdWx0PTFlLTUpOyBwLmFkZF9hcmd1bWVudCgiLS13ZWlnaHQtZGVjYXkiLHR5cGU9ZmxvYXQsZGVmYXVsdD0xZS00KQogICAgcC5hZGRfYXJndW1lbnQoIi0tbnVtLXdvcmtlcnMiLHR5cGU9aW50LGRlZmF1bHQ9MikKICAgIGE9cC5wYXJzZV9hcmdzKCkKCiAgICBjbGFzc2VzPWxvYWRfY2xhc3NlcyhhLmNsYXNzZXMpCiAgICBpZDJsYWJlbD17aTpuIGZvciBpLG4gaW4gZW51bWVyYXRlKGNsYXNzZXMpfQogICAgbGFiZWwyaWQ9e246aSBmb3IgaSxuIGluIGVudW1lcmF0ZShjbGFzc2VzKX0KCiAgICBwcm9jPVJURGV0ckltYWdlUHJvY2Vzc29yLmZyb21fcHJldHJhaW5lZChCQVNFX01PREVMKQogICAgdHJhaW49Q09DT0RldGVjdGlvbkRhdGFzZXQoUGF0aChhLnRyYWluX2RpcikvImltYWdlcyIsUGF0aChhLnRyYWluX2RpcikvImFubm90YXRpb25zLmpzb24iLHByb2MpCiAgICB2YWw9Q09DT0RldGVjdGlvbkRhdGFzZXQoUGF0aChhLnZhbF9kaXIpLyJpbWFnZXMiLFBhdGgoYS52YWxfZGlyKS8iYW5ub3RhdGlvbnMuanNvbiIscHJvYykKCiAgICBpZiBsZW4odHJhaW4pPT0wIG9yIGxlbih2YWwpPT0wOgogICAgICAgIHJhaXNlIFZhbHVlRXJyb3IoIlRyYWluaW5nIGFuZCB2YWxpZGF0aW9uIGRhdGFzZXRzIG11c3QgY29udGFpbiBhdCBsZWFzdCBvbmUgaW1hZ2UuIikKICAgIGlmIGxlbih0cmFpbi5jYXRlZ29yeV9pZF90b19sYWJlbCkhPWxlbihjbGFzc2VzKSBvciBsZW4odmFsLmNhdGVnb3J5X2lkX3RvX2xhYmVsKSE9bGVuKGNsYXNzZXMpOgogICAgICAgIHJhaXNlIFZhbHVlRXJyb3IoIkNPQ08gY2F0ZWdvcmllcyBkbyBub3QgbWF0Y2ggY2xhc3Nlcy50eHQuIFJlYnVpbGQgdGhlIGRhdGFzZXQgYWZ0ZXIgc2F2aW5nIHRoZSBjbGFzc2VzLiIpCgogICAgbW9kZWw9UlREZXRyRm9yT2JqZWN0RGV0ZWN0aW9uLmZyb21fcHJldHJhaW5lZCgKICAgICAgICBCQVNFX01PREVMLG51bV9sYWJlbHM9bGVuKGNsYXNzZXMpLGlkMmxhYmVsPWlkMmxhYmVsLGxhYmVsMmlkPWxhYmVsMmlkLAogICAgICAgIGlnbm9yZV9taXNtYXRjaGVkX3NpemVzPVRydWUKICAgICkKICAgIGRldmljZT10b3JjaC5kZXZpY2UoImN1ZGEiIGlmIHRvcmNoLmN1ZGEuaXNfYXZhaWxhYmxlKCkgZWxzZSAiY3B1IikKICAgIG1vZGVsLnRvKGRldmljZSkKICAgIHBhdGNoX3J0ZGV0cl9kZW5vaXNpbmdfZGV2aWNlKCkKCiAgICB0cj1EYXRhTG9hZGVyKHRyYWluLGJhdGNoX3NpemU9YS5iYXRjaF9zaXplLHNodWZmbGU9VHJ1ZSxudW1fd29ya2Vycz0wLGNvbGxhdGVfZm49Y29sbGF0ZV9mbikKICAgIHZhPURhdGFMb2FkZXIodmFsLGJhdGNoX3NpemU9YS5iYXRjaF9zaXplLHNodWZmbGU9RmFsc2UsbnVtX3dvcmtlcnM9MCxjb2xsYXRlX2ZuPWNvbGxhdGVfZm4pCiAgICBvcHQ9dG9yY2gub3B0aW0uQWRhbVcobW9kZWwucGFyYW1ldGVycygpLGxyPWEubGVhcm5pbmdfcmF0ZSx3ZWlnaHRfZGVjYXk9YS53ZWlnaHRfZGVjYXkpCgogICAgb3V0ZGlyPVBhdGgoYS5vdXRwdXRfZGlyKTsgb3V0ZGlyLm1rZGlyKHBhcmVudHM9VHJ1ZSxleGlzdF9vaz1UcnVlKQogICAgYmVzdD1mbG9hdCgiaW5mIikKCiAgICBmb3IgZXBvY2ggaW4gcmFuZ2UoYS5lcG9jaHMpOgogICAgICAgIG1vZGVsLnRyYWluKCk7IHJ1bm5pbmc9MAogICAgICAgIGJhcj10cWRtKHRyLGRlc2M9ZiJlcG9jaCB7ZXBvY2grMX0ve2EuZXBvY2hzfSIpCiAgICAgICAgZm9yIHN0ZXAsYmF0Y2ggaW4gZW51bWVyYXRlKGJhcik6CiAgICAgICAgICAgIGJhdGNoPW1vdmVfdG9fZGV2aWNlKGJhdGNoLCBkZXZpY2UpCiAgICAgICAgICAgICMgUlQtREVUUidzIGxvc3MgbWF0Y2hlciB1c2VzIG5lc3RlZCB0YXJnZXQgdGVuc29ycyAoYm94ZXMvY2xhc3NlcykuCiAgICAgICAgICAgICMgTW92ZSBldmVyeSB0ZW5zb3IgaW4gbGFiZWxzIHRvIHRoZSBzYW1lIGRldmljZSBhcyB0aGUgbW9kZWwuCiAgICAgICAgICAgIGlmICJsYWJlbHMiIGluIGJhdGNoOgogICAgICAgICAgICAgICAgIyBSVC1ERVRSIGV4cGVjdHMgZXZlcnkgbmVzdGVkIHRhcmdldCB0ZW5zb3Igb24gdGhlIHNhbWUgZGV2aWNlIGFzIHRoZSBtb2RlbC4KICAgICAgICAgICAgICAgIGZvciB0YXJnZXQgaW4gYmF0Y2hbImxhYmVscyJdOgogICAgICAgICAgICAgICAgICAgIGlmIGlzaW5zdGFuY2UodGFyZ2V0LCBkaWN0KToKICAgICAgICAgICAgICAgICAgICAgICAgZm9yIGtleSwgdmFsdWUgaW4gbGlzdCh0YXJnZXQuaXRlbXMoKSk6CiAgICAgICAgICAgICAgICAgICAgICAgICAgICBpZiB0b3JjaC5pc190ZW5zb3IodmFsdWUpOgogICAgICAgICAgICAgICAgICAgICAgICAgICAgICAgIHRhcmdldFtrZXldID0gdmFsdWUudG8oZGV2aWNlKQogICAgICAgICAgICBsb3NzPW1vZGVsKCoqYmF0Y2gpLmxvc3MKICAgICAgICAgICAgbG9zcy5iYWNrd2FyZCgpOyBvcHQuc3RlcCgpOyBvcHQuemVyb19ncmFkKHNldF90b19ub25lPVRydWUpCiAgICAgICAgICAgIHJ1bm5pbmcgKz0gZmxvYXQobG9zcy5pdGVtKCkpCiAgICAgICAgICAgIGJhci5zZXRfcG9zdGZpeChsb3NzPWYie3J1bm5pbmcvKHN0ZXArMSk6LjRmfSIpCiAgICAgICAgdmw9ZXZhbHVhdGUobW9kZWwsdmEsZGV2aWNlKQogICAgICAgIHByaW50KGYidmFsaWRhdGlvbl9sb3NzPXt2bDouNGZ9IikKICAgICAgICBpZiB2bDxiZXN0OgogICAgICAgICAgICBiZXN0PXZsCiAgICAgICAgICAgIG1vZGVsLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpCiAgICAgICAgICAgIHByb2Muc2F2ZV9wcmV0cmFpbmVkKG91dGRpcikKICAgICAgICAgICAgKG91dGRpci8iY2xhc3Nlcy5qc29uIikud3JpdGVfdGV4dChqc29uLmR1bXBzKHsiaWQybGFiZWwiOmlkMmxhYmVsLCJsYWJlbDJpZCI6bGFiZWwyaWR9LGluZGVudD0yKSkKICAgIG1vZGVsLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpOyBwcm9jLnNhdmVfcHJldHJhaW5lZChvdXRkaXIpCgppZiBfX25hbWVfXz09Il9fbWFpbl9fIjogbWFpbigpCg==" | |
| TRAINING_DIR = ROOT / "training" | |
| TRAIN_SCRIPT = TRAINING_DIR / "train.py" | |
| if not TRAIN_SCRIPT.exists(): | |
| TRAINING_DIR.mkdir(parents=True, exist_ok=True) | |
| TRAIN_SCRIPT.write_bytes(__import__("base64").b64decode(_TRAINING_SCRIPT_B64)) | |
| IMAGE_DIR = BASE / "images" | |
| DATASET_FILE = BASE / "dataset.json" | |
| MODEL_DIR = BASE / "model" | |
| GENERATED_DIR = BASE / "generated_dataset" | |
| CLASSES_FILE = BASE / "classes.txt" | |
| IMAGE_DIR.mkdir(parents=True, exist_ok=True) | |
| BASE.mkdir(parents=True, exist_ok=True) | |
| MODEL_DIR.mkdir(parents=True, exist_ok=True) | |
| DEFAULT_CLASSES = [ | |
| "Carnavalita", "Kimo-COno", "Squizz", "Oreo", "Moro", | |
| "Dulce", "KitKat", "Cadbury", "Mega", "other" | |
| ] | |
| if not CLASSES_FILE.exists(): | |
| CLASSES_FILE.write_text("\n".join(DEFAULT_CLASSES) + "\n", encoding="utf-8") | |
| CONFIDENCE_THRESHOLD = float(os.getenv("CONFIDENCE_THRESHOLD", "0.35")) | |
| MAX_IMAGE_MB = int(os.getenv("MAX_IMAGE_MB", "15")) | |
| _training = {"running": False, "message": "not started", "error": None} | |
| _model = None | |
| _processor = None | |
| _model_lock = threading.Lock() | |
| _annotation_click = None | |
| def load_dataset(): | |
| if not DATASET_FILE.exists(): | |
| return {"images": [], "classes": read_classes()} | |
| try: | |
| data = json.loads(DATASET_FILE.read_text(encoding="utf-8")) | |
| data.setdefault("images", []) | |
| data["classes"] = read_classes() | |
| return data | |
| except Exception: | |
| return {"images": [], "classes": read_classes()} | |
| def save_dataset(data): | |
| data["classes"] = read_classes() | |
| tmp = DATASET_FILE.with_suffix(".tmp") | |
| tmp.write_text(json.dumps(data, indent=2, ensure_ascii=False), encoding="utf-8") | |
| tmp.replace(DATASET_FILE) | |
| def read_classes(): | |
| if not CLASSES_FILE.exists(): | |
| return [] | |
| return [x.strip() for x in CLASSES_FILE.read_text(encoding="utf-8").splitlines() if x.strip()] | |
| def image_path(image_id): | |
| return IMAGE_DIR / f"{image_id}.jpg" | |
| def model_ready(): | |
| return (MODEL_DIR / "config.json").exists() | |
| def load_model(): | |
| global _model, _processor | |
| if not model_ready(): | |
| raise RuntimeError("No trained model yet. Train the model first.") | |
| with _model_lock: | |
| if _model is None: | |
| _processor = RTDetrImageProcessor.from_pretrained(str(MODEL_DIR)) | |
| _model = RTDetrForObjectDetection.from_pretrained(str(MODEL_DIR)) | |
| _model.to("cuda" if torch.cuda.is_available() else "cpu") | |
| _model.eval() | |
| return _processor, _model | |
| def dataset_status(): | |
| data = load_dataset() | |
| annotated = sum(bool(x.get("annotations")) for x in data["images"]) | |
| return ( | |
| f"**Dataset:** {len(data['images'])} images | " | |
| f"**Annotated:** {annotated} | " | |
| f"**Classes:** {len(read_classes())} | " | |
| f"**Model:** {'READY' if model_ready() else 'NOT TRAINED'} | " | |
| f"**Storage:** `{BASE}`" | |
| ) | |
| def image_choices(): | |
| data = load_dataset() | |
| return [(x["filename"], x["id"]) for x in data["images"]] | |
| def upload_training_images(files): | |
| if not files: | |
| return dataset_status(), gr.update(choices=image_choices()), "No files selected." | |
| data = load_dataset() | |
| saved = 0 | |
| skipped = [] | |
| for f in files: | |
| try: | |
| # Gradio 6 may return FileData objects or plain dictionaries. | |
| if isinstance(f, dict): | |
| raw_path = f.get("path") or f.get("name") or f.get("filepath") | |
| else: | |
| raw_path = getattr(f, "path", None) or getattr(f, "name", None) or f | |
| path = Path(raw_path) | |
| raw = path.read_bytes() | |
| if len(raw) > MAX_IMAGE_MB * 1024 * 1024: | |
| skipped.append(f"{path.name}: over {MAX_IMAGE_MB} MB") | |
| continue | |
| im = Image.open(io.BytesIO(raw)).convert("RGB") | |
| image_id = __import__("uuid").uuid4().hex | |
| out = image_path(image_id) | |
| im.save(out, "JPEG", quality=95) | |
| data["images"].append({ | |
| "id": image_id, | |
| "filename": path.name, | |
| "width": im.width, | |
| "height": im.height, | |
| "annotations": [], | |
| }) | |
| saved += 1 | |
| except Exception as e: | |
| skipped.append(f"{path.name}: {e}") | |
| save_dataset(data) | |
| msg = f"Saved {saved} image(s)." | |
| if skipped: | |
| msg += "\nSkipped:\n- " + "\n- ".join(skipped) | |
| return dataset_status(), gr.update(choices=image_choices()), msg | |
| def image_data_uri(image_id): | |
| import base64 | |
| p = image_path(image_id) | |
| if not p.exists(): | |
| return "" | |
| return "data:image/jpeg;base64," + base64.b64encode(p.read_bytes()).decode("ascii") | |
| def annotation_canvas_html(image_id): | |
| if not image_id: | |
| return '<div class="anno-empty">Select an image from the Dataset tab.</div>' | |
| data = load_dataset() | |
| item = next((x for x in data["images"] if x["id"] == image_id), None) | |
| if not item: | |
| return '<div class="anno-empty">Image not found.</div>' | |
| src = image_data_uri(image_id) | |
| boxes = json.dumps(item.get("annotations", []), ensure_ascii=False) | |
| return """<div class="anno-wrap"> | |
| <div class="anno-toolbar"><b>Draw boxes directly on the image</b><span>Click + drag + release = create box</span><span>Choose the class first</span></div> | |
| <div class="anno-canvas-wrap"><canvas id="anno-canvas"></canvas></div> | |
| <div class="anno-help" id="anno-hint">Drag from one corner of the object to the opposite corner, then release. The coordinates are filled automatically; click <b>Save Box</b>. Repeat for every object.</div> | |
| </div> | |
| <script> | |
| (function(){ | |
| const imgSrc=%s, imageW=%d, imageH=%d, saved=%s; | |
| const canvas=document.getElementById('anno-canvas'); if(!canvas)return; | |
| const ctx=canvas.getContext('2d'), img=new Image(); let drawing=false,sx=0,sy=0,current=null; | |
| function fit(){const maxW=Math.min(1100,window.innerWidth-80),maxH=Math.max(400,window.innerHeight*.62),scale=Math.min(maxW/imageW,maxH/imageH,1);canvas.width=Math.max(1,Math.round(imageW*scale));canvas.height=Math.max(1,Math.round(imageH*scale));canvas.dataset.scale=scale;redraw();} | |
| function redraw(){if(!img.complete)return;const sc=+canvas.dataset.scale||1;ctx.clearRect(0,0,canvas.width,canvas.height);ctx.drawImage(img,0,0,canvas.width,canvas.height);saved.forEach((a,i)=>{const b=a.box||[];const x=b[0]*sc,y=b[1]*sc,w=b[2]*sc,h=b[3]*sc;ctx.strokeStyle='#ff3030';ctx.lineWidth=3;ctx.strokeRect(x,y,w,h);ctx.fillStyle='#ff3030';ctx.fillRect(x,Math.max(0,y-22),120,22);ctx.fillStyle='#fff';ctx.font='14px sans-serif';ctx.fillText((i+1)+'. '+a.class,x+5,Math.max(16,y-6));});if(current){ctx.strokeStyle='#00ff88';ctx.lineWidth=3;ctx.setLineDash([7,5]);ctx.strokeRect(current.x,current.y,current.w,current.h);ctx.setLineDash([]);}} | |
| function pos(e){const r=canvas.getBoundingClientRect();return{x:e.clientX-r.left,y:e.clientY-r.top};} | |
| function setField(id,val){const box=document.querySelector('#'+id);const el=box?.querySelector('input,textarea');if(!el)return;const setter=Object.getOwnPropertyDescriptor(HTMLInputElement.prototype,'value')?.set||Object.getOwnPropertyDescriptor(HTMLTextAreaElement.prototype,'value')?.set;if(setter)setter.call(el,String(val));else el.value=String(val);el.dispatchEvent(new Event('input',{bubbles:true}));el.dispatchEvent(new Event('change',{bubbles:true}));} | |
| canvas.addEventListener('pointerdown',e=>{e.preventDefault();canvas.setPointerCapture(e.pointerId);const p=pos(e);sx=p.x;sy=p.y;drawing=true;current={x:sx,y:sy,w:0,h:0};redraw();}); | |
| canvas.addEventListener('pointermove',e=>{if(!drawing)return;const p=pos(e);current={x:Math.min(sx,p.x),y:Math.min(sy,p.y),w:Math.abs(p.x-sx),h:Math.abs(p.y-sy)};redraw();}); | |
| canvas.addEventListener('pointerup',e=>{if(!drawing)return;drawing=false;const p=pos(e),sc=+canvas.dataset.scale||1;const x=Math.min(sx,p.x)/sc,y=Math.min(sy,p.y)/sc,w=Math.abs(p.x-sx)/sc,h=Math.abs(p.y-sy)/sc;current=null;redraw();if(w<3||h<3)return;setField('anno-x',Math.round(x));setField('anno-y',Math.round(y));setField('anno-w',Math.round(w));setField('anno-h',Math.round(h));const hint=document.getElementById('anno-hint');if(hint)hint.textContent='Box created. Click “Save Box” to store it.';}); | |
| canvas.addEventListener('pointercancel',()=>{drawing=false;current=null;redraw();}); | |
| img.onload=fit;img.src=imgSrc;window.addEventListener('resize',fit); | |
| })();</script>""" % (json.dumps(src), int(item["width"]), int(item["height"]), boxes) | |
| def annotation_preview_image(image_id): | |
| """Return the real PIL image for Gradio's Image component.""" | |
| if not image_id: | |
| return None | |
| p = image_path(image_id) | |
| if not p.exists(): | |
| return None | |
| try: | |
| return Image.open(p).convert("RGB") | |
| except Exception: | |
| return None | |
| def annotation_preview_with_boxes(image_id): | |
| image = annotation_preview_image(image_id) | |
| if image is None: | |
| return None | |
| data = load_dataset() | |
| item = next((x for x in data["images"] if x["id"] == image_id), None) | |
| if not item: | |
| return image | |
| out = image.copy() | |
| draw = ImageDraw.Draw(out) | |
| for i, a in enumerate(item.get("annotations", []), 1): | |
| x, y, w, h = a["box"] | |
| draw.rectangle([x, y, x+w, y+h], outline="red", width=5) | |
| label = f"{i}. {a['class']}" | |
| y0 = max(0, y-24) | |
| draw.rectangle([x, y0, x+max(130, len(label)*9), y0+24], fill="red") | |
| draw.text((x+4, y0+4), label, fill="white") | |
| return out | |
| def refresh_editor(image_id): | |
| if not image_id: | |
| return None, "Select an image.", [] | |
| data=load_dataset() | |
| item=next((x for x in data["images"] if x["id"]==image_id),None) | |
| if not item: | |
| return None,"Image not found.",[] | |
| return annotation_preview_with_boxes(image_id), f"**{item['filename']}** — {item['width']} × {item['height']} px", item.get("annotations",[]) | |
| def draw_annotations(image_id): | |
| """Render the annotation canvas for the selected image.""" | |
| return annotation_canvas_html(image_id) | |
| def add_annotation(image_id, cls, x, y, w, h): | |
| if not image_id:return annotation_preview_with_boxes(image_id),"Select an image first.",[] | |
| if not cls:return annotation_preview_with_boxes(image_id),"Select a class first.",[] | |
| try:x,y,w,h=map(float,[x,y,w,h]) | |
| except:return annotation_preview_with_boxes(image_id),"Enter box coordinates first.",[] | |
| if w<=0 or h<=0:return annotation_preview_with_boxes(image_id),"Box must have a width and height.",[] | |
| data=load_dataset() | |
| item=next((z for z in data["images"] if z["id"]==image_id),None) | |
| if not item:return None,"Image not found.",[] | |
| x=max(0,min(x,item["width"]-1)); y=max(0,min(y,item["height"]-1)) | |
| w=min(w,item["width"]-x); h=min(h,item["height"]-y) | |
| item.setdefault("annotations",[]).append({"class":cls,"box":[x,y,w,h]}) | |
| save_dataset(data) | |
| return annotation_preview_with_boxes(image_id),f"Saved {cls}: [{x:.0f}, {y:.0f}, {w:.0f}, {h:.0f}]",item["annotations"] | |
| def remove_annotation(image_id,index): | |
| if not image_id:return annotation_preview_with_boxes(image_id),"Select an image first.",[] | |
| data=load_dataset() | |
| item=next((z for z in data["images"] if z["id"]==image_id),None) | |
| if not item:return None,"Image not found.",[] | |
| try:idx=int(index)-1 | |
| except:return annotation_preview_with_boxes(image_id),"Enter an annotation number.",item.get("annotations",[]) | |
| anns=item.get("annotations",[]) | |
| if idx<0 or idx>=len(anns):return annotation_preview_with_boxes(image_id),"Annotation number not found.",anns | |
| deleted=anns.pop(idx);save_dataset(data) | |
| return annotation_preview_with_boxes(image_id),f"Deleted annotation {index}: {deleted['class']}",anns | |
| def clear_annotations(image_id): | |
| if not image_id:return annotation_preview_with_boxes(image_id),"Select an image first.",[] | |
| data=load_dataset() | |
| item=next((z for z in data["images"] if z["id"]==image_id),None) | |
| if not item:return None,"Image not found.",[] | |
| item["annotations"]=[];save_dataset(data) | |
| return annotation_preview_with_boxes(image_id),"Annotations cleared.",[] | |
| def save_classes(text): | |
| classes = [x.strip() for x in (text or "").splitlines() if x.strip()] | |
| if not classes: | |
| return "At least one class is required.", gr.update(choices=read_classes()), dataset_status() | |
| if len(set(classes)) != len(classes): | |
| return "Classes must be unique.", gr.update(choices=read_classes()), dataset_status() | |
| CLASSES_FILE.write_text("\n".join(classes) + "\n", encoding="utf-8") | |
| data = load_dataset() | |
| save_dataset(data) | |
| return f"Saved {len(classes)} classes.", gr.update(choices=classes, value=classes[0]), dataset_status() | |
| def handle_annotation_click(image_id, cls, click_state, evt: gr.SelectData): | |
| """Use two clicks on the real Gradio image to define a box. | |
| First click = top-left corner, second click = opposite corner. | |
| This avoids the unreliable HTML canvas/script path and works in Gradio itself. | |
| """ | |
| if not image_id: | |
| return 0, 0, 0, 0, [], "Select an image first." | |
| if not cls: | |
| return 0, 0, 0, 0, [], "Select a class first." | |
| data = load_dataset() | |
| item = next((x for x in data["images"] if x["id"] == image_id), None) | |
| if not item: | |
| return 0, 0, 0, 0, [], "Image not found." | |
| try: | |
| point = evt.index | |
| px, py = float(point[0]), float(point[1]) | |
| except Exception: | |
| return 0, 0, 0, 0, click_state or [], "Could not read the image click position." | |
| px = max(0, min(px, item["width"] - 1)) | |
| py = max(0, min(py, item["height"] - 1)) | |
| state = list(click_state or []) | |
| if not state: | |
| return round(px), round(py), 0, 0, [px, py], f"First corner: ({px:.0f}, {py:.0f}). Now click the opposite corner." | |
| x0, y0 = state[:2] | |
| x = min(x0, px); y = min(y0, py) | |
| w = abs(px - x0); h = abs(py - y0) | |
| if w < 2 or h < 2: | |
| return round(x), round(y), 0, 0, [], "Box is too small. Click the first corner again." | |
| return round(x), round(y), round(w), round(h), [], f"Box ready: [{x:.0f}, {y:.0f}, {w:.0f}, {h:.0f}] for {cls}. Click Save Box." | |
| def build_coco(): | |
| data = load_dataset() | |
| classes = read_classes() | |
| if not classes: | |
| raise RuntimeError("No classes configured.") | |
| items = [x for x in data["images"] if x.get("annotations")] | |
| if not items: | |
| raise RuntimeError("Annotate at least 1 image before training.") | |
| # With only one annotated image, use it for both training and validation so | |
| # the first training run is possible. With 2+ images, use an 80/20 split. | |
| if len(items) == 1: | |
| train_items, val_items = items, items | |
| else: | |
| split = max(1, int(len(items) * 0.8)) | |
| if split >= len(items): | |
| split = len(items) - 1 | |
| train_items, val_items = items[:split], items[split:] | |
| category_id = {name: i + 1 for i, name in enumerate(classes)} | |
| def make_coco(selected): | |
| images, annotations = [], [] | |
| ann_id = 1 | |
| for item in selected: | |
| images.append({ | |
| "id": item["id"], | |
| "file_name": item["id"] + ".jpg", | |
| "width": item["width"], | |
| "height": item["height"], | |
| }) | |
| for ann in item["annotations"]: | |
| x, y, w, h = ann["box"] | |
| annotations.append({ | |
| "id": ann_id, | |
| "image_id": item["id"], | |
| "category_id": category_id[ann["class"]], | |
| "bbox": [x, y, w, h], | |
| "area": w*h, | |
| "iscrowd": 0, | |
| }) | |
| ann_id += 1 | |
| return { | |
| "images": images, | |
| "annotations": annotations, | |
| "categories": [{"id": i+1, "name": n} for i, n in enumerate(classes)] | |
| } | |
| if GENERATED_DIR.exists(): | |
| shutil.rmtree(GENERATED_DIR) | |
| for name, selected in [("train", train_items), ("val", val_items)]: | |
| d = GENERATED_DIR / name | |
| (d / "images").mkdir(parents=True, exist_ok=True) | |
| for item in selected: | |
| shutil.copy2(image_path(item["id"]), d / "images" / f"{item['id']}.jpg") | |
| (d / "annotations.json").write_text( | |
| json.dumps(make_coco(selected), indent=2), encoding="utf-8" | |
| ) | |
| def run_training(epochs, batch_size, learning_rate): | |
| global _training, _model | |
| try: | |
| _training = {"running": True, "message": "building COCO dataset", "error": None} | |
| build_coco() | |
| _training["message"] = "training RT-DETR" | |
| cmd = [ | |
| sys.executable, str(TRAIN_SCRIPT), | |
| "--train-dir", str(GENERATED_DIR / "train"), | |
| "--val-dir", str(GENERATED_DIR / "val"), | |
| "--classes", str(CLASSES_FILE), | |
| "--output-dir", str(MODEL_DIR), | |
| "--epochs", str(int(epochs)), | |
| "--batch-size", str(int(batch_size)), | |
| "--learning-rate", str(float(learning_rate)), | |
| ] | |
| result = subprocess.run(cmd, cwd=ROOT, capture_output=True, text=True) | |
| if result.returncode != 0: | |
| details = result.stderr.strip() or result.stdout.strip() or f"training process exited with code {result.returncode}" | |
| raise RuntimeError(details[-12000:]) | |
| _model = None | |
| _training = {"running": False, "message": "training complete", "error": None} | |
| except Exception as e: | |
| _training = {"running": False, "message": "training failed", "error": str(e)} | |
| def start_training(epochs, batch_size, learning_rate): | |
| """Start training from the Gradio event itself. | |
| Calling the @spaces.GPU function directly is important on Hugging Face | |
| ZeroGPU: starting it from a normal Python background thread can bypass the | |
| GPU allocation context, making the button appear to do nothing. | |
| """ | |
| global _training | |
| if _training["running"]: | |
| return json.dumps(_training, indent=2) | |
| # Validate parameters before requesting GPU time. | |
| try: | |
| epochs = max(1, int(epochs)) | |
| batch_size = max(1, int(batch_size)) | |
| learning_rate = float(learning_rate) | |
| if learning_rate <= 0: | |
| raise ValueError("Learning rate must be greater than 0.") | |
| except Exception as e: | |
| _training = {"running": False, "message": "training failed", "error": f"Invalid training settings: {e}"} | |
| return json.dumps(_training, indent=2) | |
| annotated = sum(bool(x.get("annotations")) for x in load_dataset()["images"]) | |
| if annotated < 1: | |
| _training = {"running": False, "message": "training failed", "error": "Annotate at least 1 image before training."} | |
| return json.dumps(_training, indent=2) | |
| run_training(epochs, batch_size, learning_rate) | |
| return json.dumps(_training, indent=2) | |
| def training_status(): | |
| return json.dumps(_training, indent=2) | |
| def count_image(image): | |
| if image is None: | |
| return None, "Upload an image first.", {} | |
| if not model_ready(): | |
| return None, "Model is not trained yet. Go to Training.", {} | |
| try: | |
| proc, detector = load_model() | |
| image = image.convert("RGB") if isinstance(image, Image.Image) else Image.fromarray(image).convert("RGB") | |
| device = next(detector.parameters()).device | |
| inputs = proc(images=image, return_tensors="pt") | |
| inputs = {k: v.to(device) if torch.is_tensor(v) else v for k, v in inputs.items()} | |
| with torch.inference_mode(): | |
| outputs = detector(**inputs) | |
| target_sizes = torch.tensor([[image.height, image.width]], device=device) | |
| result = proc.post_process_object_detection( | |
| outputs, threshold=CONFIDENCE_THRESHOLD, target_sizes=target_sizes | |
| )[0] | |
| detections = [] | |
| counts = Counter() | |
| for score, label, box in zip(result["scores"], result["labels"], result["boxes"]): | |
| s = float(score.item()) | |
| cls = detector.config.id2label[int(label.item())] | |
| coords = [round(float(v), 2) for v in box.tolist()] | |
| detections.append({"class": cls, "confidence": round(s, 4), "box": coords}) | |
| counts[cls] += 1 | |
| out = image.copy() | |
| draw = ImageDraw.Draw(out) | |
| for d in detections: | |
| x1, y1, x2, y2 = d["box"] | |
| draw.rectangle([x1, y1, x2, y2], outline="red", width=4) | |
| label = f"{d['class']} {d['confidence']:.2f}" | |
| draw.rectangle([x1, max(0, y1-22), x1+max(120, len(label)*8), y1], fill="red") | |
| draw.text((x1+3, max(0, y1-20)), label, fill="white") | |
| response = { | |
| "total": len(detections), | |
| "counts": dict(sorted(counts.items())), | |
| "detections": detections, | |
| } | |
| return out, json.dumps(response, indent=2), response["counts"] | |
| except Exception as e: | |
| return None, f"Counting failed: {e}", {} | |
| # ----- Gradio UI ----- | |
| CSS = """ | |
| .gradio-container { max-width: 1250px !important; } | |
| h1 { margin-bottom: 0.2rem !important; } | |
| .anno-wrap{width:100%}.anno-toolbar{display:flex;gap:14px;flex-wrap:wrap;padding:10px 12px;margin-bottom:8px;border-radius:10px;background:#20242a}.anno-toolbar span{opacity:.85}.anno-canvas-wrap{width:100%;overflow:auto;border:1px solid #555;border-radius:10px;background:#111;padding:8px}.anno-canvas-wrap canvas{display:block;max-width:none;cursor:crosshair;touch-action:none;margin:auto}.anno-help{padding:8px 2px;opacity:.75}.anno-empty{padding:50px;text-align:center;border:1px dashed #777;border-radius:10px} | |
| #annotation-image img { max-height: 650px !important; object-fit: contain !important; } | |
| .status { padding: 10px 14px; border-radius: 10px; } | |
| """ | |
| with gr.Blocks(title="Ice Cream Dataset + Counter") as demo: | |
| gr.Markdown("# 🍦 Ice Cream Dataset + Counter\nUpload and annotate training images, train RT-DETR, then count ice creams in new images.") | |
| status = gr.Markdown(dataset_status(), elem_classes="status") | |
| with gr.Tab("1 · Dataset"): | |
| gr.Markdown("### Upload training images") | |
| files = gr.Files(file_count="multiple", file_types=["image"], type="filepath", label="Images") | |
| upload_btn = gr.Button("Save Images", variant="primary") | |
| upload_msg = gr.Markdown() | |
| # Dataset selector | |
| image_select = gr.Dropdown(choices=image_choices(), label="Training image", interactive=True) | |
| refresh_btn = gr.Button("Refresh Dataset") | |
| refresh_btn.click(lambda: (dataset_status(), gr.update(choices=image_choices())), None, [status, image_select]) | |
| gr.Markdown("### Classes") | |
| class_text = gr.Textbox(value="\n".join(read_classes()), lines=8, label="One class per line") | |
| save_class_btn = gr.Button("Save Classes") | |
| class_msg = gr.Markdown() | |
| # Class dropdown is updated after the Annotate tab creates it. | |
| with gr.Tab("2 · Annotate"): | |
| gr.Markdown("### Annotate training images") | |
| gr.Markdown("Select an image above, choose a class, then **click the first corner and click the opposite corner** of each object. The real uploaded image is shown below. Click **Save Box** after each box.") | |
| with gr.Row(): | |
| with gr.Column(scale=3): | |
| annotation_image = gr.Image(value=None, type="pil", interactive=False, label="Training image", height=650, elem_id="annotation-image") | |
| editor_info = gr.Markdown("Select an image from the Dataset tab.") | |
| with gr.Column(scale=1): | |
| ann_class = gr.Dropdown(choices=read_classes(), value=(read_classes()[0] if read_classes() else None), label="Class", interactive=True) | |
| x = gr.Number(label="X (left)", value=0, precision=0) | |
| y = gr.Number(label="Y (top)", value=0, precision=0) | |
| w = gr.Number(label="Width", value=0, precision=0) | |
| h = gr.Number(label="Height", value=0, precision=0) | |
| add_btn = gr.Button("💾 Save Box", variant="primary") | |
| gr.Markdown("**Box method:** click corner 1 → click corner 2 → Save Box.") | |
| delete_index = gr.Number(label="Annotation # to delete", value=1, precision=0) | |
| delete_btn = gr.Button("Delete Box") | |
| clear_btn = gr.Button("Clear All Boxes") | |
| annotations = gr.JSON(label="Saved annotations") | |
| ann_msg = gr.Markdown() | |
| click_state = gr.State([]) | |
| def load_annotation_image(image_id): | |
| return annotation_preview_with_boxes(image_id), refresh_editor(image_id)[1], refresh_editor(image_id)[2], [] | |
| image_select.change(load_annotation_image, image_select, [annotation_image, editor_info, annotations, click_state]) | |
| annotation_image.select(handle_annotation_click, [image_select, ann_class, click_state], [x, y, w, h, click_state, ann_msg]) | |
| add_btn.click(add_annotation, [image_select, ann_class, x, y, w, h], [annotation_image, ann_msg, annotations]) | |
| delete_btn.click(remove_annotation, [image_select, delete_index], [annotation_image, ann_msg, annotations]) | |
| clear_btn.click(clear_annotations, image_select, [annotation_image, ann_msg, annotations]) | |
| with gr.Tab("3 · Training"): | |
| gr.Markdown("### Train RT-DETR") | |
| gr.Markdown("Training runs in the Space process. A GPU Space is strongly recommended for practical training speed.") | |
| with gr.Row(): | |
| epochs = gr.Number(value=int(os.getenv("EPOCHS", "30")), label="Epochs", precision=0) | |
| batch = gr.Number(value=int(os.getenv("BATCH_SIZE", "2")), label="Batch size", precision=0) | |
| lr = gr.Number(value=float(os.getenv("LEARNING_RATE", "1e-5")), label="Learning rate") | |
| train_btn = gr.Button("🚀 Start Training", variant="primary") | |
| refresh_train = gr.Button("Refresh Training Status") | |
| train_out = gr.Code(value=training_status, language="json", label="Training status") | |
| train_btn.click(start_training, [epochs, batch, lr], train_out) | |
| refresh_train.click(training_status, None, train_out) | |
| with gr.Tab("4 · Count"): | |
| gr.Markdown("### Count ice creams") | |
| count_in = gr.Image(type="pil", sources=["upload", "clipboard"], label="Image to count") | |
| count_btn = gr.Button("🍦 Count", variant="primary") | |
| count_out = gr.Image(label="Detections") | |
| count_json = gr.Code(language="json", label="Detection details") | |
| count_table = gr.JSON(label="Counts by class") | |
| count_btn.click(count_image, count_in, [count_out, count_json, count_table]) | |
| # Correct the upload event now that image_select exists. | |
| upload_btn.click( | |
| upload_training_images, | |
| files, | |
| [status, image_select, upload_msg], | |
| preprocess=False, | |
| queue=False, | |
| ) | |
| save_class_btn.click(save_classes, class_text, [class_msg, ann_class, status]) | |
| # The earlier placeholder event is harmlessly superseded by this real event. | |
| demo.load(lambda: (dataset_status(), gr.update(choices=image_choices()), gr.update(choices=read_classes(), value=(read_classes()[0] if read_classes() else None))), | |
| None, [status, image_select, ann_class]) | |
| if __name__ == "__main__": | |
| demo.queue().launch( | |
| server_name="0.0.0.0", | |
| server_port=int(os.getenv("PORT", "7860")), | |
| css=CSS, | |
| ) | |