model_train_dm/backend/apps/datasets/views/annotate_task_sample.py

587 lines
20 KiB
Python
Raw Permalink Normal View History

2026-07-27 17:51:49 +08:00
import os
import uuid
from datetime import datetime
import xml.etree.ElementTree as ET
from django.db import transaction
from django.db.models import Case, Count, IntegerField, Q, Sum, When
from django.shortcuts import get_object_or_404
from rest_framework import status
from rest_framework.decorators import api_view, permission_classes
from rest_framework.pagination import PageNumberPagination
from rest_framework.permissions import IsAuthenticated
from rest_framework.response import Response
from rest_framework.utils import json
from apps.common.serializers import AiAnnotateTaskSampleSerializer
from apps.common.utils.annotation_store import (
get_dataset_sample_annotation_content,
split_dataset_sample_annotation_content,
)
from apps.core.models import (
AiAnnotateTask,
AiAnnotateTaskSample,
AiDatasetSample,
)
from apps.datasets.dataset_stats import refresh_dataset_sample_counters
from apps.datasets.storage_service import download_sidecar_text, upload_sidecar_bytes
def _safe_split_wh(custom1):
try:
w_str, h_str = str(custom1).lower().split("x")
return int(w_str), int(h_str)
except Exception:
return None, None
def _extract_sample_index(saved_filename):
stem = os.path.splitext(os.path.basename(str(saved_filename or "")))[0]
tail = stem.rsplit("-", 1)[-1]
if tail.isdigit():
return int(tail)
return None
def _get_dataset_sample_summary_status(dataset_sample, task_sample_statuses):
if dataset_sample.rectangle or dataset_sample.polygon or "06" in task_sample_statuses:
return "02"
if task_sample_statuses:
return "01"
return "00"
def _refresh_dataset_sample_status(sample_ids):
if not sample_ids:
return
sample_ids = list({sample_id for sample_id in sample_ids if sample_id})
samples = list(
AiDatasetSample.objects.filter(id__in=sample_ids).only("id", "status", "rectangle", "polygon")
)
if not samples:
return
task_sample_status_map = {}
for sample_id, status_code in AiAnnotateTaskSample.objects.filter(sample_id__in=sample_ids).values_list(
"sample_id", "status"
):
task_sample_status_map.setdefault(sample_id, set()).add(status_code)
objects_to_update = []
for sample in samples:
new_status = _get_dataset_sample_summary_status(
sample, task_sample_status_map.get(sample.id, set())
)
if sample.status != new_status:
objects_to_update.append(AiDatasetSample(id=sample.id, status=new_status))
if objects_to_update:
AiDatasetSample.objects.bulk_update(objects_to_update, fields=["status"])
def _normalize_polygon_points(region):
if not isinstance(region, list):
return []
pts = []
for p in region:
if not isinstance(p, (list, tuple)) or len(p) != 2:
continue
try:
x = float(p[0])
y = float(p[1])
except Exception:
continue
pts.append([x, y])
return pts if len(pts) >= 3 else []
def _shape_key(label, points):
try:
norm = [(round(float(x), 2), round(float(y), 2)) for x, y in points]
except Exception:
return None
return str(label), json.dumps(norm, separators=(",", ":"), ensure_ascii=False)
def _is_polygon_items(items):
if not isinstance(items, list):
return False
for it in items:
if not isinstance(it, dict):
continue
region = it.get("region")
if isinstance(region, list) and region and isinstance(region[0], (list, tuple)):
return True
return False
def _bbox_key(item):
try:
label = str(item.get("label"))
region = item.get("region")
if not isinstance(region, list) or len(region) != 4:
return None
vals = [round(float(x), 2) for x in region]
return label, json.dumps(vals, separators=(",", ":"), ensure_ascii=False)
except Exception:
return None
def _shapes_from_items(items):
if not isinstance(items, list):
return []
shapes = []
for it in items:
if not isinstance(it, dict):
continue
label = it.get("label")
if not label:
continue
points = _normalize_polygon_points(it.get("region"))
if not points:
continue
shapes.append(
{
"label": str(label),
"points": points,
"group_id": None,
"shape_type": "polygon",
"flags": {},
"mask": None,
}
)
return shapes
def _shapes_from_labelme(labelme_obj):
if not isinstance(labelme_obj, dict):
return []
shapes_obj = labelme_obj.get("shapes")
if not isinstance(shapes_obj, list):
return []
shapes = []
for sh in shapes_obj:
if not isinstance(sh, dict):
continue
if sh.get("shape_type") != "polygon":
continue
label = sh.get("label")
points = _normalize_polygon_points(sh.get("points"))
if not label or not points:
continue
shapes.append(
{
"label": str(label),
"points": points,
"group_id": None,
"shape_type": "polygon",
"flags": {},
"mask": None,
}
)
return shapes
def _load_polygon_shapes_from_dataset_sample(dataset_sample):
content = get_dataset_sample_annotation_content(dataset_sample)
if not content:
return []
try:
parsed = json.loads(content)
except Exception:
return []
if isinstance(parsed, dict) and parsed.get("type") == "polygon":
try:
json_text = download_sidecar_text(dataset_sample.saved_path, dataset_sample.saved_filename, ".json")
if not json_text:
return []
labelme_obj = json.loads(json_text)
except Exception:
return []
return _shapes_from_labelme(labelme_obj)
if isinstance(parsed, list):
return _shapes_from_items(parsed)
return []
@api_view(["POST"])
@permission_classes([IsAuthenticated])
def set_person_task(request):
data = request.data
task_id = data["task_id"]
task = AiAnnotateTask.objects.get(id=task_id)
dataset_id = task.dataset_id
status_code = data.get("status")
task_range = task.task_range
query_condition = Q(dataset_id=dataset_id)
if status_code:
query_condition &= Q(status=status_code)
all_samples = list(AiDatasetSample.objects.filter(query_condition).order_by("saved_filename"))
existing_sample_ids = set(
AiAnnotateTaskSample.objects.filter(task_id=task_id).values_list("sample_id", flat=True)
)
if existing_sample_ids:
all_samples = [sample for sample in all_samples if sample.id not in existing_sample_ids]
if task_range:
start, end = map(int, task_range.split("-"))
all_samples = [
sample
for sample in all_samples
if (sample_index := _extract_sample_index(sample.saved_filename)) is not None
and start <= sample_index <= end
]
all_samples_length = len(all_samples)
rule_type = data["rule_type"]
task_rules = json.loads(data["task_rules"])
total_percent = 0
last_person = ""
for key, value in task_rules.items():
person = key
if rule_type == "percent":
percent = value
total_percent += percent
last_person = person
num_samples_to_assign = int(all_samples_length * percent / 100)
create_person_samples(task_id, person, all_samples[:num_samples_to_assign])
all_samples = all_samples[num_samples_to_assign:]
elif rule_type == "number":
number = value
create_person_samples(task_id, person, all_samples[:number])
all_samples = all_samples[number:]
elif rule_type == "range":
start, end = map(int, value.split("-"))
matching_samples = [
sample
for sample in all_samples
if (sample_index := _extract_sample_index(sample.saved_filename)) is not None
and start <= sample_index <= end
]
create_person_samples(task_id, person, matching_samples)
all_samples = list(set(all_samples) - set(matching_samples))
if rule_type == "percent":
remaining_samples = len(all_samples)
if total_percent == 100 and remaining_samples > 0:
create_person_samples(task_id, last_person, all_samples)
return Response({"message": "人员标注任务分配成功"})
@transaction.atomic
def create_person_samples(task_id, person, assigned_samples):
if not assigned_samples:
return
objects_to_save = [
AiAnnotateTaskSample(
id=str(uuid.uuid4()),
task_id=task_id,
executor=person,
sample_id=sample.id,
sample_path=sample.saved_path,
sample_name=sample.saved_filename,
status="01",
custom1=sample.custom1,
)
for sample in assigned_samples
]
AiAnnotateTaskSample.objects.bulk_create(objects_to_save)
_refresh_dataset_sample_status([sample.id for sample in assigned_samples])
@api_view(["POST"])
@permission_classes([IsAuthenticated])
def annotate_task_sample(request):
data = request.data
sample_id = data.get("id")
hard = data.get("hard")
task_sample = get_object_or_404(AiAnnotateTaskSample, id=sample_id)
AiAnnotateTaskSample.objects.filter(id=sample_id).update(
annotation_content=data.get("annotation_content"),
status="02",
hard=hard,
annotation_time=datetime.now(),
)
_refresh_dataset_sample_status([task_sample.sample_id])
return Response({"message": "ok"})
@api_view(["POST"])
@permission_classes([IsAuthenticated])
def submit_task_samples(request):
data = request.data
task_id = data.get("task_id")
person = data.get("person")
ids = [id for id in data.get("ids", "").split(",") if id]
queryset = AiAnnotateTaskSample.objects.filter(task_id=task_id, executor=person, status="02")
if ids:
queryset = queryset.filter(id__in=ids)
sample_ids = list(queryset.values_list("sample_id", flat=True))
queryset.update(status="03")
_refresh_dataset_sample_status(sample_ids)
return Response({"message": "ok"})
@api_view(["GET"])
@permission_classes([IsAuthenticated])
def read_task_sample(request):
sample_id = request.GET.get("id", None)
sample = get_object_or_404(AiAnnotateTaskSample, id=sample_id)
serializer = AiAnnotateTaskSampleSerializer(sample)
return Response(serializer.data)
@api_view(["POST"])
@permission_classes([IsAuthenticated])
def audit_task_samples(request):
data = request.data
task_id = data.get("task_id")
person = request.user.username
feedback = data.get("feedback")
status_code = data.get("status")
ids = [id for id in data.get("ids", "").split(",") if id]
queryset = AiAnnotateTaskSample.objects.filter(task_id=task_id, status="03")
if ids:
queryset = queryset.filter(id__in=ids)
if status_code == "04":
sample_ids = list(queryset.values_list("sample_id", flat=True))
queryset.update(status="04", auditor=person)
_refresh_dataset_sample_status(sample_ids)
elif status_code == "05":
sample_ids = list(queryset.values_list("sample_id", flat=True))
queryset.update(status="05", auditor=person, feedback=feedback)
_refresh_dataset_sample_status(sample_ids)
return Response({"message": "ok"})
@api_view(["POST"])
@permission_classes([IsAuthenticated])
@transaction.atomic
def merge_sample_content(request):
try:
data = request.data
task_id = data.get("task_id")
if not task_id:
return Response({"error": "缺少 task_id"}, status=status.HTTP_400_BAD_REQUEST)
datasettask = AiAnnotateTask.objects.get(id=task_id)
executor = data.get("executor", None)
filter_kwargs = {"task_id": task_id}
if executor:
filter_kwargs["executor"] = executor
filter_kwargs["status"] = "04"
queryset = AiAnnotateTaskSample.objects.filter(**filter_kwargs)
samples = list(queryset)
sample_ids = [sample.sample_id for sample in samples]
with transaction.atomic():
for sample in samples:
if sample.annotation_content:
person_items = json.loads(sample.annotation_content)
dataset_sample = AiDatasetSample.objects.get(id=sample.sample_id)
if datasettask.annotate_type == "03" or _is_polygon_items(person_items):
annotation_fields = split_dataset_sample_annotation_content(sample.annotation_content)
AiDatasetSample.objects.filter(id=sample.sample_id).update(
rectangle=annotation_fields["rectangle"],
polygon=annotation_fields["polygon"],
)
else:
dataset_annotation = get_dataset_sample_annotation_content(dataset_sample)
dataset_items = json.loads(dataset_annotation) if dataset_annotation else []
merged_data = person_items + dataset_items
unique_annotations = {}
for item in merged_data:
key = _bbox_key(item)
if not key:
continue
if key not in unique_annotations:
unique_annotations[key] = item
final_merged_data = list(unique_annotations.values())
imgwidth, imgheight = _safe_split_wh(sample.custom1)
xml_bytes = _build_voc_xml_bytes(
final_merged_data,
sample.sample_path,
sample.sample_name,
imgwidth,
imgheight,
)
upload_sidecar_bytes(
sample.sample_path,
sample.sample_name,
".xml",
xml_bytes,
"application/xml",
)
merged_json = json.dumps(final_merged_data, ensure_ascii=False)
annotation_fields = split_dataset_sample_annotation_content(merged_json)
AiDatasetSample.objects.filter(id=sample.sample_id).update(
rectangle=annotation_fields["rectangle"],
polygon=annotation_fields["polygon"],
)
queryset.update(status="06")
_refresh_dataset_sample_status(sample_ids)
refresh_dataset_sample_counters(datasettask.dataset_id)
sample_count = AiAnnotateTaskSample.objects.filter(task_id=task_id).count()
audited_count = AiAnnotateTaskSample.objects.filter(
task_id=task_id, status__in=["04", "06"]
).count()
if sample_count == audited_count:
AiAnnotateTask.objects.filter(id=task_id).update(status="03")
return Response({"message": "任务样本成功合并到数据集中"})
except Exception as e:
return Response(
{"error": "合并入库失败", "details": str(e)},
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
def _build_voc_xml_bytes(data, sample_path, filename, width, height):
normalized_path = str(sample_path or "").strip("/").replace("\\", "/")
image_path = f"{normalized_path}/{filename}" if normalized_path else filename
annotation = ET.Element("annotation")
folder = ET.SubElement(annotation, "folder")
folder.text = normalized_path
filename_elem = ET.SubElement(annotation, "filename")
filename_elem.text = filename
path = ET.SubElement(annotation, "path")
path.text = image_path
source = ET.SubElement(annotation, "source")
database = ET.SubElement(source, "database")
database.text = "Unknown"
size = ET.SubElement(annotation, "size")
width_elem = ET.SubElement(size, "width")
width_elem.text = str(int(width) if width is not None else 0)
height_elem = ET.SubElement(size, "height")
height_elem.text = str(int(height) if height is not None else 0)
depth = ET.SubElement(size, "depth")
depth.text = "3"
segmented = ET.SubElement(annotation, "segmented")
segmented.text = "0"
for obj in data:
object_elem = ET.SubElement(annotation, "object")
name = ET.SubElement(object_elem, "name")
name.text = obj["label"]
pose = ET.SubElement(object_elem, "pose")
pose.text = "Unspecified"
truncated = ET.SubElement(object_elem, "truncated")
truncated.text = "0"
difficult = ET.SubElement(object_elem, "difficult")
difficult.text = "0"
bndbox = ET.SubElement(object_elem, "bndbox")
xmin, ymin, xmax, ymax = map(float, obj["region"])
xmin_elem = ET.SubElement(bndbox, "xmin")
xmin_elem.text = str(int(xmin))
ymin_elem = ET.SubElement(bndbox, "ymin")
ymin_elem.text = str(int(ymin))
xmax_elem = ET.SubElement(bndbox, "xmax")
xmax_elem.text = str(int(xmax))
ymax_elem = ET.SubElement(bndbox, "ymax")
ymax_elem.text = str(int(ymax))
return ET.tostring(annotation, encoding="utf-8", xml_declaration=True)
@api_view(["DELETE"])
@permission_classes([IsAuthenticated])
@transaction.atomic
def delete_task_samples(request):
data = request.data
task_id = data.get("task_id")
person = data.get("person")
ids = [id for id in data.get("ids", "").split(",") if id]
try:
queryset = AiAnnotateTaskSample.objects.filter(task_id=task_id, executor=person)
if ids:
queryset = queryset.filter(id__in=ids)
samples = list(queryset)
except AiAnnotateTaskSample.DoesNotExist:
return Response({"message": "没有选择对应的样本数据!"}, status=status.HTTP_404_NOT_FOUND)
with transaction.atomic():
sample_ids = [sample.sample_id for sample in samples]
queryset.delete()
_refresh_dataset_sample_status(sample_ids)
return Response({"message": "任务样本删除成功"})
@api_view(["GET"])
@permission_classes([IsAuthenticated])
def search_task_samples(request):
task_id = request.GET.get("task_id", None)
executor = request.GET.get("executor", None)
status_code = request.GET.get("status", "")
status_in = status_code.split(",") if status_code else []
sample_name = request.GET.get("sample_name", None)
is_hard = request.GET.get("hard", None)
page_size = request.GET.get("page_size", 10)
filter_kwargs = {}
if task_id:
filter_kwargs["task_id"] = task_id
if executor:
filter_kwargs["executor"] = executor
if status_code:
filter_kwargs["status__in"] = status_in
if sample_name:
filter_kwargs["sample_name__contains"] = sample_name
if is_hard:
filter_kwargs["hard"] = is_hard
samples = AiAnnotateTaskSample.objects.filter(**filter_kwargs).order_by("sample_name")
paginator = PageNumberPagination()
paginator.page_size = page_size
result_page = paginator.paginate_queryset(samples, request)
formatted_samples = []
for sample in result_page:
sample.annotation_time = (
sample.annotation_time.strftime("%Y-%m-%d %H:%M:%S")
if sample.annotation_time
else None
)
formatted_samples.append(sample)
serializer = AiAnnotateTaskSampleSerializer(formatted_samples, many=True)
return paginator.get_paginated_response(serializer.data)
@api_view(["GET"])
@permission_classes([IsAuthenticated])
def query_person_tasks(request):
task_id = request.GET.get("task_id", None)
STATUSES = ["01", "02", "03", "04", "05", "06"]
annotate_dict = {}
for status_code in STATUSES:
annotate_dict[f"count_{status_code}"] = Sum(
Case(When(status=status_code, then=1), default=0, output_field=IntegerField())
)
annotate_dict["total"] = Count("id")
executor_status_counts = (
AiAnnotateTaskSample.objects.values("executor")
.annotate(**annotate_dict)
.filter(task_id=task_id)
.order_by("executor")
)
return Response(executor_status_counts)