model_fatsusus / patch_detected.py
Jjtumarai's picture
fix: patch build_model.py cuda calls + torch.load map_location
b957821
Raw
History Blame Contribute Delete
14.9 kB
# -*- coding: utf-8 -*-
# ============================================================
# patch_detected.py — แก้ repo 2DImage2BMI ให้รันได้บน Hugging Face Space
#
# ===== ต้องการอะไร =====
# repo ต้นฉบับรันบนเครื่องเราไม่ได้ ด้วย 2 เหตุผล:
# (1) import ของ 4 ชิ้นที่ build ไม่ผ่าน (SCHP/PSP/CPM/CRFRNN)
# (2) ฮาร์ดโค้ด .cuda() ไว้ -> ถ้าเครื่องไม่มี GPU จะพังทันที
# Kaggle มี T4 เลยเจอแค่ปัญหา (1) ส่วน (2) ไม่เคยโผล่
# พอย้ายมา Space ฟรีที่รันบน CPU -> เจอปัญหา (2) เต็มๆ
#
# ===== ไฟล์นี้ทำอะไร =====
# แก้ 2 ไฟล์ในrepo ให้ใช้งานได้:
# Detected.py -> ตัด import ที่ไม่มี + เปลี่ยน .cuda() เป็น auto-detect
# modeling/affine_align.py -> เปลี่ยน .cuda() เป็น auto-detect
# รันซ้ำได้ ไม่พัง (เช็ค marker ก่อนแก้ทุกครั้ง)
#
# ===== วิธีคิด =====
# *** ไม่ฮาร์ดโค้ดเป็น 'cpu' *** แต่ทำเป็น auto-detect:
# _DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
# -> มี GPU ก็ใช้ GPU / ไม่มีก็ใช้ CPU = ใช้ได้ทุก hardware ไม่ต้องแก้อีก
#
# *** ห้ามเอา SCHP/PSP/CPM/CRFRNN กลับมา แม้จะลงได้ ***
# ของพวกนั้นมีไว้ "ลบแขนออกจาก mask" ก่อนวัดสัดส่วน
# แต่โมเดล SVR ของเราเทรนบน mask ที่ "มีแขน" อยู่แล้ว
# ถ้าตอนใช้งานไปตัดแขน -> WSR/Area/H2W เปลี่ยนหมด -> SVR เจอเลขที่ไม่เคยเห็น -> BMI มั่ว
# กฎเหล็ก: ตอนใช้ต้องเหมือนตอนเทรนเป๊ะ รวมถึงข้อบกพร่องด้วย
# ============================================================
import sys
import os
MARK = '# [cut] ' # เครื่องหมายของบรรทัดที่ถูกคอมเมนต์ทิ้ง
TAG = '# [patched]' # เครื่องหมายของบรรทัดที่ถูกแก้ (ใช้เช็คว่า patch ไปแล้วหรือยัง)
# ---------- โค้ดตรวจหา device ที่จะแทรกเข้าไป ----------
DEVICE_SNIPPET = (
f"import torch as _torch {TAG}\n"
f"_DEVICE = 'cuda' if _torch.cuda.is_available() else 'cpu' {TAG} auto-detect: มี GPU ใช้ GPU / ไม่มีใช้ CPU"
)
# ============================================================
# รายการแก้ของแต่ละไฟล์
# ============================================================
PATCHES = {
# ---------------- ไฟล์หลัก ----------------
'Detected.py': {
# (1) คอมเมนต์ทิ้ง: import ของที่เราไม่มี + โค้ดที่เรียกใช้มัน
'cut': [
'from Human_Parse import HumanParser',
'from PSP import HumanParser_PSP',
'from CPM import CPM_Keypoint',
'from CRFRNN import CRFRNN_Contour',
'self._HumanParser = HumanParser()',
'Arms_mask = self._HumanParser.Arms_detect(img)',
'ContourOutput = ContourOutput ^ Arms_mask',
],
# (2) แทรกโค้ดตรวจ device ต่อท้าย import สุดท้าย
'insert_after': 'import time',
# (3) เปลี่ยน .cuda() เป็น .to(_DEVICE)
# ⚠️ ห้ามเติมคอมเมนต์ต่อท้ายข้อความที่แทน ถ้าจุดนั้นอยู่ "กลางนิพจน์"
# เพราะคอมเมนต์จะกลืนส่วนที่เหลือของบรรทัด -> วงเล็บไม่ปิด -> SyntaxError
# (เคยพลาดมาแล้วตอนเทส) ใส่ได้เฉพาะบรรทัดที่จบในตัวเอง
'replace': [
('Model = Pose2Seg().cuda()',
'Model = Pose2Seg().to(_DEVICE)'),
# detectron2 default เป็น cuda -> ต้องสั่งให้ตรง device ด้วย
# (บรรทัดนี้เติม tag ได้ เพราะเป็นบรรทัดเต็มที่จบในตัวเอง)
('cfg.MODEL.WEIGHTS = model_zoo.get_checkpoint_url(key_file)',
'cfg.MODEL.WEIGHTS = model_zoo.get_checkpoint_url(key_file)\n'
f' cfg.MODEL.DEVICE = _DEVICE {TAG}'),
],
},
# ---------------- ไส้ของ Pose2Seg: ตัวโมเดลหลัก ----------------
# เจอทีหลัง (7 ส.ค. รอบ 2) — ตอนแรก patch แค่ 2 ไฟล์ แล้วยังพัง
# "Found no NVIDIA driver" เพราะไฟล์นี้ยังเรียก .cuda(0) และ torch.load แบบไม่บอก device
'modeling/build_model.py': {
'cut': [],
'insert_after': 'import torch.nn.functional as F',
'replace': [
# .cuda(0) มี 5 จุด (บรรทัด 47, 50, 178, 191, 218) — replace ทีเดียวได้หมด
('.cuda(0)', '.to(_DEVICE)'),
# ===== หัวใจ: ตัวที่ทำให้ error จริง =====
# pose2seg_release.pkl ถูกเซฟจากเครื่องที่มี GPU
# torch.load แบบไม่ใส่ map_location จะพยายามคืนค่าลง GPU เดิม
# -> เครื่องไม่มี GPU = "Found no NVIDIA driver"
('pretrained_dict = torch.load(path)',
'pretrained_dict = torch.load(path, map_location=_DEVICE)'),
],
},
# ---------------- ไส้ของ Pose2Seg: การจัดตำแหน่งภาพ ----------------
'modeling/affine_align.py': {
'cut': [],
'insert_after': 'import torch.nn.functional as F',
# ทั้งสองจุดอยู่กลางนิพจน์ -> ห้ามเติมคอมเมนต์ต่อท้าย
'replace': [
('torch.from_numpy(Hs_new).cuda()',
'torch.from_numpy(Hs_new).to(_DEVICE)'),
('align_corners=True).float().cuda()',
'align_corners=True).float().to(_DEVICE)'),
],
},
}
def patch_file(path, spec):
"""
แก้ไฟล์ 1 ไฟล์ตาม spec
คืน (จำนวนที่แก้, list ของสิ่งที่ยังแก้ไม่ได้)
"""
src = open(path, encoding='utf-8').read()
n = 0
# ---------- (1) คอมเมนต์ทิ้งบรรทัดที่ไม่ต้องการ ----------
for bad in spec['cut']:
if bad in src and MARK + bad not in src: # เงื่อนไขที่ 2 = กัน patch ซ้ำ
src = src.replace(bad, MARK + bad)
n += 1
print(f' ตัด : {bad[:58]}')
# ---------- (2) แทรกโค้ดตรวจ device ----------
anchor = spec['insert_after']
if TAG not in src: # ยังไม่เคยแทรก
if anchor not in src:
return n, [f'หา anchor "{anchor}" ไม่เจอ — แทรก _DEVICE ไม่ได้']
src = src.replace(anchor, anchor + '\n' + DEVICE_SNIPPET, 1) # แทรกครั้งเดียวพอ
n += 1
print(f' แทรก : _DEVICE (auto-detect) ต่อจาก "{anchor}"')
# ---------- (3) เปลี่ยน .cuda() -> .to(_DEVICE) ----------
for old, new in spec['replace']:
if old in src and new not in src:
src = src.replace(old, new)
n += 1
print(f' เปลี่ยน: {old[:50]}')
open(path, 'w', encoding='utf-8').write(src)
# ---------- ตรวจซ้ำว่าไม่เหลืออะไรที่จะทำให้พัง ----------
left = []
for bad in spec['cut']:
if bad in src and MARK + bad not in src:
left.append(f'ยังไม่ถูกตัด: {bad}')
# ตรวจจาก "ข้อความใหม่ต้องมีอยู่จริง" ไม่ใช่ "ข้อความเก่าต้องหายไป"
# เพราะบางเคสข้อความเก่าเป็นส่วนหนึ่งของข้อความใหม่ (เช่น cfg.MODEL.WEIGHTS)
# ถ้าเช็คว่าเก่าหายไป จะฟ้องผิดทั้งที่ patch สำเร็จแล้ว
for _, new in spec['replace']:
if new not in src:
left.append(f'ยังไม่ถูกเปลี่ยนเป็น: {new.splitlines()[0][:50]}')
# ตรวจขั้นสุดท้าย: ต้องไม่เหลือ .cuda() ในไฟล์ที่อยู่ในเส้นทางรันจริง
for i, line in enumerate(src.splitlines(), 1):
if '.cuda()' in line and not line.strip().startswith('#'):
left.append(f'ยังเหลือ .cuda() บรรทัด {i}: {line.strip()[:50]}')
# ตรวจว่าไฟล์ยัง compile ผ่าน (กันพลาดแบบเติมคอมเมนต์กลางนิพจน์)
try:
compile(src, path, 'exec')
except SyntaxError as e:
left.append(f'SyntaxError บรรทัด {e.lineno}: {e.msg}')
return n, left
def main(repo):
# รับได้ทั้ง path ของ repo และ path ของ Detected.py (เผื่อเรียกแบบเดิม)
if repo.endswith('.py'):
repo = os.path.dirname(repo)
if not os.path.isdir(repo):
print(f'❌ ไม่เจอโฟลเดอร์ repo: {repo}')
return 1
print(f'patch repo: {repo}\n')
total, problems = 0, []
for rel, spec in PATCHES.items():
p = os.path.join(repo, rel)
print(f' [{rel}]')
if not os.path.exists(p):
print(f' ❌ ไม่เจอไฟล์')
problems.append(f'ไม่เจอ {rel}')
continue
n, left = patch_file(p, spec)
total += n
problems += [f'{rel}: {x}' for x in left]
print(f' -> แก้ไป {n} จุด' if n else ' -> ไม่มีอะไรต้องแก้ (patch ไปแล้ว)')
print()
if problems:
print('❌ มีปัญหา:')
for x in problems:
print(f' - {x}')
return 1
print(f'✅ patch สำเร็จ รวม {total} จุด')
# ============================================================
# ตรวจรอบสุดท้าย: สแกน "ทั้ง repo" หา .cuda ที่ยังเหลือ
#
# ทำไมต้องมี: ตอนแรกผม patch แค่ Detected.py + affine_align.py
# แล้วมั่นใจว่าครบ -> ที่จริง modeling/build_model.py ยังมี .cuda(0) อีก 5 จุด
# -> deploy ไปแล้วเจอ "Found no NVIDIA driver" เสียเวลาอีกรอบ
# การสแกนทั้ง repo ทำให้เห็นทุกจุดตั้งแต่ตอน patch ไม่ต้องรอไปพังบนเซิร์ฟเวอร์
#
# ข้ามโค้ดใต้ if __name__ == '__main__' เพราะไม่ถูกรันตอน import
# ============================================================
print('\n --- สแกนทั้ง repo หา .cuda ที่ยังเหลือ ---')
leftover = []
for dirpath, dirnames, files in os.walk(repo):
dirnames[:] = [d for d in dirnames if d not in ('__pycache__', '.git')]
for fn in files:
if not fn.endswith('.py'):
continue
p = os.path.join(dirpath, fn)
rel = os.path.relpath(p, repo).replace('\\', '/')
try:
lines = open(p, encoding='utf-8').read().splitlines()
except Exception:
continue
in_main = False
for i, line in enumerate(lines, 1):
if line.startswith("if __name__"):
in_main = True # ตั้งแต่บรรทัดนี้ลงไปไม่ถูกรันตอน import
# ข้าม: โค้ดใน __main__ / คอมเมนต์ / บรรทัดที่เรา patch เข้าไปเอง
if in_main or line.strip().startswith('#') or TAG in line:
continue
if '.cuda' in line:
leftover.append(f'{rel}:{i} {line.strip()[:60]}')
if leftover:
print(' ⚠️ ยังเหลือ .cuda (เช็คว่าไฟล์นี้ถูก import ตอนรันจริงไหม):')
for x in leftover:
print(f' {x}')
print(' หมายเหตุ: ถ้าไฟล์นั้นไม่มีใคร import ก็ปล่อยได้ ไม่ต้องแก้')
else:
print(' ✅ ไม่เหลือ .cuda ที่ไหนเลย')
print('\n✅ ตรวจแล้ว: ไม่เหลือ import ที่จะพัง และ .cuda ในไฟล์ที่ patch ถูกแก้ครบ')
return 0
if __name__ == '__main__':
target = sys.argv[1] if len(sys.argv) > 1 else '/home/user/app/2DImage2BMI-main'
sys.exit(main(target))