Skip to content

Commit 66f49c3

Browse files
committed
Revert "Add YAML creation from STAR file for training"
This reverts commit afe7134.
1 parent 6bd29a6 commit 66f49c3

2 files changed

Lines changed: 39 additions & 293 deletions

File tree

partinet/__init__.py

Lines changed: 6 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -60,19 +60,13 @@ def main():
6060
pass
6161

6262
@main.command()
63-
@click.option("--star", type=str, required=True, help="Path to input STAR file (or labels directory if using --split-only)")
64-
@click.option("--images", type=str, required=True, help="Path to directory containing micrograph images")
65-
@click.option("--output", type=str, required=True, help="Path to output directory for organized train/val data")
66-
@click.option("--class-id", type=int, default=0, help="Class ID to assign to all particles (default: 0)")
67-
@click.option("--test-size", type=float, default=0.25, help="Proportion of dataset to use for validation (default: 0.25)")
68-
@click.option("--split-only", is_flag=True, help="Skip STAR conversion and only split existing labels")
69-
def split(star, images, output, class_id, test_size, split_only):
70-
if split_only:
71-
click.echo("Splitting existing labels into training and validation sets...")
72-
else:
73-
click.echo("Converting STAR file to YOLO format and splitting for training...")
63+
@click.option("--labels", type=str, required=True, help="Path to the labels directory")
64+
@click.option("--images", type=str, required=True, help="Path to the images directory")
65+
@click.option("--output", type=str, required=True, help="Path to the output directory")
66+
def split(labels, images, output):
67+
click.echo("Splitting micrographs for training and validation...")
7468
import partinet.process_utils.split_train
75-
partinet.process_utils.split_train.main(star, images, output, class_id, test_size, split_only)
69+
partinet.process_utils.split_train.main(labels, images, output)
7670

7771
@main.command()
7872
@click.option("--labels", type=str, required=True, help="Path to the labels directory")
Lines changed: 33 additions & 281 deletions
Original file line numberDiff line numberDiff line change
@@ -1,245 +1,67 @@
11
import os
22
import shutil
33
import argparse
4-
import cv2
4+
55
import numpy as np
6-
from typing import Dict, List
76
from sklearn.model_selection import train_test_split as tts
87

9-
10-
def parse_star_file(star_path: str) -> List[Dict[str, str]]:
11-
"""
12-
Parse a STAR file and extract particle coordinates
13-
Returns a list of dictionaries with micrograph names and coordinates
14-
"""
15-
particles = []
16-
in_data_section = False
17-
headers = []
18-
header_indices = {}
19-
20-
with open(star_path, 'r') as f:
21-
for line in f:
22-
line = line.strip()
23-
24-
# Skip empty lines and comments
25-
if not line or line.startswith('#'):
26-
continue
27-
28-
# Check for data section
29-
if line.startswith('data_'):
30-
in_data_section = True
31-
continue
32-
33-
# Check for loop section
34-
if line.startswith('loop_'):
35-
continue
36-
37-
# Parse headers
38-
if line.startswith('_rln'):
39-
parts = line.split()
40-
header_name = parts[0]
41-
if len(parts) > 1:
42-
header_idx = int(parts[1].replace('#', '')) - 1 # Convert to 0-indexed
43-
headers.append(header_name)
44-
header_indices[header_name] = header_idx
45-
continue
46-
47-
# Parse data rows
48-
if in_data_section and headers:
49-
parts = line.split()
50-
if len(parts) >= len(headers):
51-
particle = {}
52-
for header in headers:
53-
idx = header_indices[header]
54-
particle[header] = parts[idx]
55-
particles.append(particle)
56-
57-
return particles
58-
59-
60-
def starfile_to_yolo(x_coord: int, y_coord: int, diameter: int,
61-
image_width: int, image_height: int, class_id: int = 0) -> Dict[str, float]:
62-
"""
63-
Convert STAR file coordinates to YOLO format
64-
"""
65-
# Calculate YOLO normalized coordinates
66-
x_center = x_coord / image_width
67-
y_center = y_coord / image_height
68-
width = diameter / image_width
69-
height = diameter / image_height
70-
71-
return {
72-
'class': class_id,
73-
'x_center': x_center,
74-
'y_center': y_center,
75-
'width': width,
76-
'height': height
77-
}
78-
79-
80-
def group_particles_by_micrograph(particles: List[Dict[str, str]]) -> Dict[str, List[Dict[str, str]]]:
81-
"""
82-
Group particles by their micrograph name
83-
"""
84-
grouped = {}
85-
for particle in particles:
86-
micrograph = particle.get('_rlnMicrographName', '')
87-
if micrograph not in grouped:
88-
grouped[micrograph] = []
89-
grouped[micrograph].append(particle)
90-
return grouped
91-
92-
93-
def convert_star_to_yolo(star_path: str, images_path: str, output_labels_path: str, class_id: int = 0) -> List[str]:
94-
"""
95-
Convert STAR file to YOLO format labels
96-
97-
Args:
98-
star_path: Path to input STAR file
99-
images_path: Path to directory containing micrograph images
100-
output_labels_path: Path to output directory for YOLO label files
101-
class_id: Class ID to assign to all particles (default: 0)
102-
103-
Returns:
104-
List of processed filenames (without extensions)
105-
"""
106-
# Create output directory if it doesn't exist
107-
os.makedirs(output_labels_path, exist_ok=True)
108-
109-
# Parse STAR file
110-
print(f"Parsing STAR file: {star_path}")
111-
particles = parse_star_file(star_path)
112-
print(f"Found {len(particles)} particles")
113-
114-
# Group particles by micrograph
115-
grouped_particles = group_particles_by_micrograph(particles)
116-
print(f"Found {len(grouped_particles)} unique micrographs")
117-
118-
processed_files = []
119-
120-
# Process each micrograph
121-
for micrograph_name, micrograph_particles in grouped_particles.items():
122-
# Remove .mrc extension and get base filename
123-
base_filename = os.path.splitext(os.path.basename(micrograph_name))[0]
124-
125-
# Find corresponding image file (try common extensions)
126-
image_file = None
127-
image_ext = None
128-
for ext in ['.png', '.jpg', '.jpeg', '.tif', '.tiff', '.mrc']:
129-
potential_path = os.path.join(images_path, base_filename + ext)
130-
if os.path.exists(potential_path):
131-
image_file = potential_path
132-
image_ext = ext
133-
break
134-
135-
if image_file is None:
136-
print(f"Warning: Image file not found for {base_filename}. Skipping.")
137-
continue
138-
139-
# Read image to get dimensions
140-
image = cv2.imread(image_file)
141-
if image is None:
142-
print(f"Warning: Could not read image {image_file}. Skipping.")
143-
continue
144-
145-
img_height, img_width = image.shape[:2]
146-
147-
# Convert particles to YOLO format
148-
yolo_labels = []
149-
for particle in micrograph_particles:
150-
x_coord = int(particle.get('_rlnCoordinateX', 0))
151-
y_coord = int(particle.get('_rlnCoordinateY', 0))
152-
diameter = int(particle.get('_rlnDiameter', 0))
153-
154-
yolo_coords = starfile_to_yolo(x_coord, y_coord, diameter,
155-
img_width, img_height, class_id)
156-
yolo_labels.append(yolo_coords)
157-
158-
# Write YOLO label file
159-
output_file = os.path.join(output_labels_path, f"{base_filename}.txt")
160-
with open(output_file, 'w') as f:
161-
for label in yolo_labels:
162-
f.write(f"{label['class']} {label['x_center']:.6f} {label['y_center']:.6f} "
163-
f"{label['width']:.6f} {label['height']:.6f}\n")
164-
165-
processed_files.append((base_filename, image_ext))
166-
print(f"Processed {base_filename}: {len(yolo_labels)} particles")
167-
168-
print(f"\nConversion complete! Labels saved to {output_labels_path}")
169-
return processed_files
170-
171-
172-
def split_train_val(labels_path: str, images_path: str, output_dir: str, test_size: float = 0.25):
8+
def main(labels_path: str, images_path: str, output_dir: str):
1739
"""
17410
Splits a dataset of images and labels into training and validation sets and organizes them into
17511
a specified output directory. Generates corresponding .txt files for train/val data and a
17612
cryo_training.yaml file for use in model training.
17713
17814
Args:
179-
labels_path: Path to the directory containing label files
180-
images_path: Path to the directory containing image files
181-
output_dir: Path to the output directory where split data will be saved
182-
test_size: Proportion of dataset to use for validation (default: 0.25)
15+
labels_path (str): Path to the directory containing label files.
16+
images_path (str): Path to the directory containing image files.
17+
output_dir (str): Path to the output directory where split data will be saved.
18318
"""
18419
# Create output directories if they do not exist
18520
if not os.path.exists(output_dir):
18621
os.makedirs(output_dir)
187-
os.makedirs(os.path.join(output_dir, "images", "train"), exist_ok=True)
188-
os.makedirs(os.path.join(output_dir, "images", "val"), exist_ok=True)
189-
os.makedirs(os.path.join(output_dir, "labels", "train"), exist_ok=True)
190-
os.makedirs(os.path.join(output_dir, "labels", "val"), exist_ok=True)
22+
os.makedirs(os.path.join(output_dir, "images", "train"))
23+
os.makedirs(os.path.join(output_dir, "images", "val"))
24+
os.makedirs(os.path.join(output_dir, "labels", "train"))
25+
os.makedirs(os.path.join(output_dir, "labels", "val"))
19126

19227
# List all label files
193-
files = [f for f in os.listdir(labels_path) if f.endswith('.txt')]
28+
files = os.listdir(labels_path)
19429

19530
# Split data into training and validation indices
196-
train_idx, val_idx = tts(np.arange(0, len(files), 1), test_size=test_size, shuffle=True, random_state=42)
31+
train_idx, val_idx = tts(np.arange(0, len(files), 1), shuffle=True)
19732

19833
# Iterate through files and copy them into train/val directories
19934
for idx, file in enumerate(files):
200-
file_name = os.path.splitext(file)[0]
201-
202-
# Find the image file (try common extensions)
203-
image_file = None
204-
for ext in ['.png', '.jpg', '.jpeg', '.tif', '.tiff', '.mrc']:
205-
potential_image = os.path.join(images_path, file_name + ext)
206-
if os.path.exists(potential_image):
207-
image_file = file_name + ext
208-
break
209-
210-
if image_file is None:
211-
print(f"Warning: Image file not found for {file_name}. Skipping.")
212-
continue
213-
35+
file_name = file[:-4] # Remove file extension
21436
if idx in train_idx:
21537
# Copy training images and labels
21638
shutil.copy(
217-
os.path.join(images_path, image_file),
218-
os.path.join(output_dir, "images", "train", image_file)
39+
os.path.join(images_path, file_name + ".png"),
40+
os.path.join(output_dir, "images", "train", file_name + ".png")
21941
)
22042
shutil.copy(
221-
os.path.join(labels_path, file),
222-
os.path.join(output_dir, "labels", "train", file)
43+
os.path.join(labels_path, file_name + ".txt"),
44+
os.path.join(output_dir, "labels", "train", file_name + ".txt")
22345
)
22446
elif idx in val_idx:
22547
# Copy validation images and labels
22648
shutil.copy(
227-
os.path.join(images_path, image_file),
228-
os.path.join(output_dir, "images", "val", image_file)
49+
os.path.join(images_path, file_name + ".png"),
50+
os.path.join(output_dir, "images", "val", file_name + ".png")
22951
)
23052
shutil.copy(
231-
os.path.join(labels_path, file),
232-
os.path.join(output_dir, "labels", "val", file)
53+
os.path.join(labels_path, file_name + ".txt"),
54+
os.path.join(output_dir, "labels", "val", file_name + ".txt")
23355
)
23456

23557
# Create val.txt file listing validation image paths
23658
with open(os.path.join(output_dir, "val.txt"), "w") as f:
237-
for file in sorted(os.listdir(os.path.join(output_dir, "images", "val"))):
59+
for file in os.listdir(os.path.join(output_dir, "images", "val")):
23860
f.write(str(os.path.join(output_dir, "images", "val", file)) + "\n")
23961

24062
# Create train.txt file listing training image paths
24163
with open(os.path.join(output_dir, "train.txt"), "w") as f:
242-
for file in sorted(os.listdir(os.path.join(output_dir, "images", "train"))):
64+
for file in os.listdir(os.path.join(output_dir, "images", "train")):
24365
f.write(str(os.path.join(output_dir, "images", "train", file)) + "\n")
24466

24567
# Create cryo_training.yaml file with dataset configuration
@@ -254,92 +76,22 @@ def split_train_val(labels_path: str, images_path: str, output_dir: str, test_si
25476

25577
with open(os.path.join(output_dir, "cryo_training.yaml"), "w") as f:
25678
f.write(to_write)
257-
258-
train_count = len(os.listdir(os.path.join(output_dir, "images", "train")))
259-
val_count = len(os.listdir(os.path.join(output_dir, "images", "val")))
260-
print(f"\nDataset split complete!")
261-
print(f"Training samples: {train_count}")
262-
print(f"Validation samples: {val_count}")
263-
print(f"Configuration saved to: {os.path.join(output_dir, 'cryo_training.yaml')}")
264-
26579

266-
def main(star_path: str, images_path: str, output_dir: str, class_id: int = 0,
267-
test_size: float = 0.25, split_only: bool = False):
268-
"""
269-
Main function to convert STAR file to YOLO format and split into train/val sets
270-
271-
Args:
272-
star_path: Path to input STAR file (or labels directory if split_only=True)
273-
images_path: Path to directory containing micrograph images
274-
output_dir: Path to output directory
275-
class_id: Class ID to assign to all particles (default: 0)
276-
test_size: Proportion of dataset to use for validation (default: 0.25)
277-
split_only: If True, skip conversion and only split existing labels
80+
def parse_args() -> argparse.Namespace:
27881
"""
279-
if split_only:
280-
# Skip conversion, just split existing labels
281-
print("Splitting existing labels into train/val sets...")
282-
split_train_val(star_path, images_path, output_dir, test_size)
283-
else:
284-
# Create temporary directory for converted labels
285-
temp_labels_dir = os.path.join(output_dir, "temp_labels")
286-
287-
# Convert STAR to YOLO
288-
print("=" * 60)
289-
print("Step 1: Converting STAR file to YOLO format")
290-
print("=" * 60)
291-
convert_star_to_yolo(star_path, images_path, temp_labels_dir, class_id)
292-
293-
# Split into train/val
294-
print("\n" + "=" * 60)
295-
print("Step 2: Splitting data into train/val sets")
296-
print("=" * 60)
297-
split_train_val(temp_labels_dir, images_path, output_dir, test_size)
298-
299-
# Clean up temporary labels directory
300-
shutil.rmtree(temp_labels_dir)
301-
print(f"\nAll done! Training data ready in {output_dir}")
82+
Parses command-line arguments for the dataset splitting script.
30283
84+
Returns:
85+
argparse.Namespace: Parsed arguments containing paths for labels, images, and output.
86+
"""
87+
parser = argparse.ArgumentParser(description="Create training data split from images and labels.")
88+
parser.add_argument("--labels", required=True, help="Path to the labels directory")
89+
parser.add_argument("--images", required=True, help="Path to the images directory")
90+
parser.add_argument("--output", required=True, help="Path to the output directory")
30391

304-
def parse_args() -> argparse.Namespace:
305-
parser = argparse.ArgumentParser(
306-
description="Convert STAR file to YOLO format and split into train/val sets"
307-
)
308-
parser.add_argument(
309-
"--star",
310-
required=True,
311-
help="Path to input STAR file (or labels directory if using --split-only)"
312-
)
313-
parser.add_argument(
314-
"--images",
315-
required=True,
316-
help="Path to directory containing micrograph images"
317-
)
318-
parser.add_argument(
319-
"--output",
320-
required=True,
321-
help="Path to output directory for organized train/val data"
322-
)
323-
parser.add_argument(
324-
"--class-id",
325-
type=int,
326-
default=0,
327-
help="Class ID to assign to all particles (default: 0)"
328-
)
329-
parser.add_argument(
330-
"--test-size",
331-
type=float,
332-
default=0.25,
333-
help="Proportion of dataset to use for validation (default: 0.25)"
334-
)
335-
parser.add_argument(
336-
"--split-only",
337-
action="store_true",
338-
help="Skip STAR conversion and only split existing labels (use --star to specify labels directory)"
339-
)
34092
return parser.parse_args()
34193

342-
34394
if __name__ == "__main__":
95+
# Parse command-line arguments and execute the main function
34496
args = parse_args()
345-
main(args.star, args.images, args.output, args.class_id, args.test_size, args.split_only)
97+
main(args.labels, args.images, args.output)

0 commit comments

Comments
 (0)