/*Copyright ©2025 APIJSON(https://github.com/APIJSON) Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except in compliance with the License. You may obtain a copy of the License at http://www.apache.org/licenses/LICENSE-2.0 Unless required by applicable law or agreed to in writing, software distributed under the License is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific language governing permissions and limitations under the License.*/ package apijson; import com.google.gson.Gson; import com.google.gson.GsonBuilder; import java.io.FileWriter; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; import java.nio.file.Paths; import java.text.SimpleDateFormat; import java.util.*; public class DatasetUtil { public static void main(String[] args) { try { // --- 调用示例 --- // 示例1:只生成检测数据集 System.out.println("Generating DETECTION dataset..."); Set detectionTasks = new HashSet<>(Collections.singletonList(TaskType.DETECTION)); generate("./output/detection_dataset", detectionTasks); // 示例2:生成分割数据集 System.out.println("\nGenerating SEGMENTATION dataset..."); Set segTasks = new HashSet<>(Collections.singletonList(TaskType.SEGMENTATION)); generate("./output/segmentation_dataset", segTasks); // 示例3:生成姿态关键点数据集 System.out.println("\nGenerating POSE_KEYPOINTS dataset..."); Set keypointTasks = new HashSet<>(Collections.singletonList(TaskType.POSE_KEYPOINTS)); generate("./output/keypoints_dataset", keypointTasks); // 示例4:生成OCR数据集 System.out.println("\nGenerating OCR dataset..."); Set ocrTasks = new HashSet<>(Collections.singletonList(TaskType.OCR)); generate("./output/ocr_dataset", ocrTasks); // 示例5:在一个JSON中同时包含检测和关键点标注 System.out.println("\nGenerating combined DETECTION and KEYPOINTS dataset..."); Set combinedTasks = new HashSet<>(Arrays.asList(TaskType.DETECTION, TaskType.POSE_KEYPOINTS)); generate("./output/combined_dataset", combinedTasks); } catch (IOException e) { e.printStackTrace(); } } /** * 定义支持的任务类型 */ public enum TaskType { CLASSIFICATION, DETECTION, SEGMENTATION, POSE_KEYPOINTS, FACE_KEYPOINTS, ROTATED_DETECTION, OCR } /** * 数据集构建器 */ public static class DatasetBuilder { private final CocoDataset dataset; private int imageIdCounter = 1; private int annotationIdCounter = 1; public DatasetBuilder() { this.dataset = new CocoDataset(); this.dataset.setInfo(new HashMap<>()); this.dataset.setLicenses(new ArrayList<>()); this.dataset.setImages(new ArrayList<>()); this.dataset.setCategories(new ArrayList<>()); this.dataset.setAnnotations(new ArrayList<>()); } public DatasetBuilder withInfo(String description, String version, String year) { Map info = new HashMap<>(); info.put("description", description); info.put("version", version); info.put("year", year); info.put("date_created", new SimpleDateFormat("yyyy-MM-dd").format(new Date())); this.dataset.setInfo(info); return this; } public DatasetBuilder withCategory(int id, String name, String supercategory) { Category cat = new Category(); cat.setId(id); cat.setName(name); cat.setSupercategory(supercategory); this.dataset.getCategories().add(cat); return this; } // 可为关键点任务添加专门的 category 方法 public DatasetBuilder withKeypointCategory(int id, String name, String supercategory, List keypoints, List> skeleton) { Category cat = new Category(); cat.setId(id); cat.setName(name); cat.setSupercategory(supercategory); cat.setKeypoints(keypoints); cat.setSkeleton(skeleton); this.dataset.getCategories().add(cat); return this; } public DatasetBuilder addImage(String fileName, int width, int height) { ImageInfo img = new ImageInfo(); img.setId(imageIdCounter++); img.setFile_name(fileName); img.setWidth(width); img.setHeight(height); this.dataset.getImages().add(img); return this; } public DatasetBuilder addAnnotation(Annotation annotation) { // 确保设置了唯一的 ID annotation.setId(annotationIdCounter++); this.dataset.getAnnotations().add(annotation); return this; } public CocoDataset build() { return this.dataset; } } /** * 将 COCO 数据集对象写入 JSON 文件 * @param dataset COCO 数据集对象 * @param outputPath 输出文件路径 (e.g., /path/to/annotations/instances_train2017.json) */ public static void writeToFile(CocoDataset dataset, String outputPath) throws IOException { Path parentDir = Paths.get(outputPath).getParent(); if (parentDir != null && !Files.exists(parentDir)) { Files.createDirectories(parentDir); } Gson gson = new GsonBuilder().setPrettyPrinting().create(); try (FileWriter writer = new FileWriter(outputPath)) { gson.toJson(dataset, writer); } System.out.println("Successfully generated COCO JSON file at: " + outputPath); } /** * 主生成方法(示例) * 实际使用中,你需要从你的数据源(如XML, CSV)读取数据来填充这些 Annotation */ public static void generate(String outputDir, Set tasks) throws IOException { // --- 1. 初始化构建器和通用信息 --- DatasetBuilder builder = new DatasetBuilder() .withInfo("My Custom Dataset", "1.0", "2025") .withCategory(1, "person", "person") .withCategory(2, "car", "vehicle") .withCategory(3, "dog", "animal") .withKeypointCategory(1, "person", "person", Arrays.asList("nose", "left_eye", "right_eye"), // 简化版关键点 Arrays.asList(Arrays.asList(1, 2), Arrays.asList(1, 3)) ); // --- 2. 添加图片信息 --- // 假设我们有两张图片 builder.addImage("00001.jpg", 640, 480); // image_id 将是 1 builder.addImage("00002.jpg", 800, 600); // image_id 将是 2 // --- 3. 根据任务类型添加标注 (核心部分) --- // 这是示例数据,你需要替换成你自己的真实数据加载逻辑 // 为 image 1 添加标注 if (tasks.contains(TaskType.DETECTION) || tasks.contains(TaskType.SEGMENTATION) || tasks.contains(TaskType.ROTATED_DETECTION)) { DetectionAnnotation detAnn = new DetectionAnnotation(); detAnn.setImage_id(1); detAnn.setCategory_id(3); // dog detAnn.setBbox(Arrays.asList(100.0, 50.0, 80.0, 120.0)); detAnn.setArea(80.0 * 120.0); if (tasks.contains(TaskType.SEGMENTATION)) { detAnn.setSegmentation(Arrays.asList( Arrays.asList(100.0, 50.0, 180.0, 50.0, 180.0, 170.0, 100.0, 170.0) )); } if(tasks.contains(TaskType.ROTATED_DETECTION)){ // 旋转检测通常用四点表示,这里也放在segmentation里 detAnn.setSegmentation(Arrays.asList( Arrays.asList(110.0, 55.0, 175.0, 60.0, 170.0, 165.0, 105.0, 160.0) )); } builder.addAnnotation(detAnn); } if (tasks.contains(TaskType.POSE_KEYPOINTS)) { KeypointAnnotation kpAnn = new KeypointAnnotation(); kpAnn.setImage_id(1); kpAnn.setCategory_id(1); // person kpAnn.setBbox(Arrays.asList(200.0, 100.0, 50.0, 150.0)); kpAnn.setArea(50.0 * 150.0); kpAnn.setNum_keypoints(3); kpAnn.setKeypoints(Arrays.asList(225.0, 110.0, 2.0, 215.0, 105.0, 2.0, 235.0, 105.0, 2.0)); // [x,y,v, x,y,v, ...] builder.addAnnotation(kpAnn); } // 为 image 2 添加标注 if (tasks.contains(TaskType.OCR)) { OcrAnnotation ocrAnn = new OcrAnnotation(); ocrAnn.setImage_id(2); ocrAnn.setCategory_id(2); // car, 假设车牌是OCR目标 ocrAnn.setBbox(Arrays.asList(300.0, 400.0, 120.0, 30.0)); ocrAnn.setArea(120.0 * 30.0); // OCR通常用四边形表示位置 ocrAnn.setSegmentation(Arrays.asList( Arrays.asList(300.0, 400.0, 420.0, 400.0, 420.0, 430.0, 300.0, 430.0) )); Map attrs = new HashMap<>(); attrs.put("transcription", "AB-1234"); attrs.put("legible", true); ocrAnn.setAttributes(attrs); builder.addAnnotation(ocrAnn); } // --- 4. 构建并写入文件 --- CocoDataset cocoDataset = builder.build(); // 为不同任务生成不同的文件名 String taskName = tasks.iterator().next().toString().toLowerCase(); // 用第一个任务命名 String outputJsonPath = Paths.get(outputDir, "annotations", "instances_" + taskName + ".json").toString(); writeToFile(cocoDataset, outputJsonPath); // 之后,你需要将图片文件(00001.jpg, 00002.jpg)复制到指定的图片目录下, // 例如 outputDir/images/ } public static class ImageInfo { private int id; private String file_name; private int width; private int height; public int getId() { return id; } public void setId(int id) { this.id = id; } public String getFile_name() { return file_name; } public void setFile_name(String file_name) { this.file_name = file_name; } public int getWidth() { return width; } public void setWidth(int width) { this.width = width; } public int getHeight() { return height; } public void setHeight(int height) { this.height = height; } } public static class Category { private int id; private String name; private String supercategory; // For Keypoints private List keypoints; private List> skeleton; public int getId() { return id; } public void setId(int id) { this.id = id; } public String getName() { return name; } public void setName(String name) { this.name = name; } public String getSupercategory() { return supercategory; } public void setSupercategory(String supercategory) { this.supercategory = supercategory; } public List getKeypoints() { return keypoints; } public void setKeypoints(List keypoints) { this.keypoints = keypoints; } public List> getSkeleton() { return skeleton; } public void setSkeleton(List> skeleton) { this.skeleton = skeleton; } } public static class Annotation { private int id; private int image_id; private int category_id; public int getId() { return id; } public void setId(int id) { this.id = id; } public int getImage_id() { return image_id; } public void setImage_id(int image_id) { this.image_id = image_id; } public int getCategory_id() { return category_id; } public void setCategory_id(int category_id) { this.category_id = category_id; } } public static class DetectionAnnotation extends Annotation { private List bbox; // [x, y, width, height] private double area; private int iscrowd = 0; private List> segmentation; // for segmentation & rotated box public List getBbox() { return bbox; } public void setBbox(List bbox) { this.bbox = bbox; } public double getArea() { return area; } public void setArea(double area) { this.area = area; } public int getIscrowd() { return iscrowd; } public void setIscrowd(int iscrowd) { this.iscrowd = iscrowd; } public List> getSegmentation() { return segmentation; } public void setSegmentation(List> segmentation) { this.segmentation = segmentation; } } public static class KeypointAnnotation extends DetectionAnnotation { private int num_keypoints; private List keypoints; // [x1, y1, v1, x2, y2, v2, ...] public int getNum_keypoints() { return num_keypoints; } public void setNum_keypoints(int num_keypoints) { this.num_keypoints = num_keypoints; } public List getKeypoints() { return keypoints; } public void setKeypoints(List keypoints) { this.keypoints = keypoints; } } public static class OcrAnnotation extends DetectionAnnotation { private Map attributes; // {"transcription": "TEXT", "legible": true} public Map getAttributes() { return attributes; } public void setAttributes(Map attributes) { this.attributes = attributes; } } public static class CocoDataset { private Map info; private List> licenses; private List images; private List categories; private List annotations; // 使用基类,实现多态 public Map getInfo() { return info; } public void setInfo(Map info) { this.info = info; } public List> getLicenses() { return licenses; } public void setLicenses(List> licenses) { this.licenses = licenses; } public List getImages() { return images; } public void setImages(List images) { this.images = images; } public List getCategories() { return categories; } public void setCategories(List categories) { this.categories = categories; } public List getAnnotations() { return annotations; } public void setAnnotations(List annotations) { this.annotations = annotations; } } }