Files
2026-07-09 18:21:21 +09:00

124 lines
3.7 KiB
Python

import os
import json
import torch
import argparse
import requests
from io import BytesIO
from PIL import Image
from torchvision import transforms as T
from net import get_model
num_cls_dict = {'market': 30, 'duke': 23}
num_ids_dict = {'market': 751, 'duke': 702}
transforms = T.Compose([
T.Resize(size=(288, 144)),
T.ToTensor(),
T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
class PredictDecoder(object):
def __init__(self, dataset):
with open('./doc/label.json', 'r') as f:
self.label_list = json.load(f)[dataset]
with open('./doc/attribute.json', 'r') as f:
self.attribute_dict = json.load(f)[dataset]
self.dataset = dataset
self.num_label = len(self.label_list)
def decode(self, pred):
pred = pred.squeeze(dim=0)
results = {}
for idx in range(self.num_label):
name, choice = self.attribute_dict[self.label_list[idx]]
value = choice[pred[idx]]
if value:
results[name] = value
return results
def load_network(network, dataset, model_name):
save_path = os.path.join('./checkpoints', dataset, model_name, 'net_last.pth')
network.load_state_dict(torch.load(save_path))
print(f'[+] 모델 로드 완료: {save_path}')
return network
def preprocess_image(image):
src = transforms(image)
return src.unsqueeze(dim=0)
def load_image_from_url_or_path(image_source):
if image_source.startswith(('http://', 'https://')):
response = requests.get(image_source)
response.raise_for_status()
return Image.open(BytesIO(response.content)).convert('RGB')
if not os.path.isfile(image_source):
raise FileNotFoundError(f"Image not found: {image_source}")
return Image.open(image_source).convert('RGB')
def parse_xywh(xywh):
parts = xywh.split(",")
if len(parts) != 4:
raise ValueError(f"xywh must be 'x,y,w,h', got: {xywh}")
return tuple(int(v) for v in parts)
def predict(image, dataset='market', backbone='resnet50', use_id=False):
assert dataset in ['market', 'duke']
assert backbone in ['resnet50', 'resnet34', 'resnet18', 'densenet121']
model_name = f'{backbone}_nfc_id' if use_id else f'{backbone}_nfc'
num_label = num_cls_dict[dataset]
num_id = num_ids_dict[dataset]
model = get_model(model_name, num_label, use_id=use_id, num_id=num_id)
model = load_network(model, dataset, model_name)
model.eval()
src = preprocess_image(image)
with torch.no_grad():
if not use_id:
out = model.forward(src)
else:
out, _ = model.forward(src)
pred = torch.gt(out, torch.ones_like(out) / 2)
decoder = PredictDecoder(dataset)
return decoder.decode(pred)
def run(image_url, xywh, dataset='market', backbone='resnet50', use_id=False):
oimg = load_image_from_url_or_path(image_url)
x, y, w, h = parse_xywh(xywh)
cimg = oimg.crop((x, y, x + w, y + h))
results = predict(cimg, dataset=dataset, backbone=backbone, use_id=use_id)
print("\n" + "=" * 50)
print(" Person 상세 속성 분석 결과 ")
print("=" * 50)
for name, value in results.items():
print(f'{name}: {value}')
print("=" * 50)
return results
if __name__ == "__main__":
# python start.py --image_url "https://acai.ketidev.kr:20443/detect/image/202606/20260619_145116_image.jpg" --xywh "404,290,74,193"
parser = argparse.ArgumentParser(description="Person Attribute Recognition Agent")
parser.add_argument("--image_url", type=str, required=True)
parser.add_argument("--xywh", type=str, required=True)
args = parser.parse_args()
run(args.image_url, args.xywh)