-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathsunburst_script.py
More file actions
862 lines (721 loc) · 37.9 KB
/
Copy pathsunburst_script.py
File metadata and controls
862 lines (721 loc) · 37.9 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
#!/usr/bin/env python3
"""
Enhanced Sunburst Chart Generator
Creates hierarchical sunburst plots from CSV data with support for 3-5 levels
of hierarchy, intelligent color inheritance, smart labeling, and publication-quality output.
Usage:
python sunburst_script.py <csv_file> [options]
Basic Examples:
# Simple 3-level chart with default settings
python sunburst_script.py data.csv
# Clean professional chart with thin borders
python sunburst_script.py data.csv --line-width 0.3 --threshold 8.0
# Count unique species instead of all records
python sunburst_script.py data.csv --sample-id "Species" --count-unique
# Custom hierarchy levels
python sunburst_script.py data.csv --level1 "Kingdom" --level2 "Phylum" --level3 "Class"
# Four or five level hierarchies
python sunburst_script.py data.csv --level4 "Order" --level5 "Family"
Grouping Examples:
# Keep top 10 items per level, aggregate rest into "Other"
python sunburst_script.py data.csv --top-n 10
# Use global threshold (percentage of total dataset)
python sunburst_script.py data.csv --threshold 2.0 --threshold-mode global
# Combined: top 8 items, then filter out anything < 1% globally
python sunburst_script.py data.csv --top-n 8 --threshold 1.0 --threshold-mode global
Label Formatting Examples:
# Show only names (no counts)
python sunburst_script.py data.csv --label-style name-only
# Show names with percentages
python sunburst_script.py data.csv --label-style name-percent
# Full labels: name, count, and percentage
python sunburst_script.py data.csv --label-style full
# Add percentage to default name-count style
python sunburst_script.py data.csv --show-percent
Advanced Examples:
# Color inheritance with variations (progressive shading)
python sunburst_script.py data.csv --color-inherit-level 1 --color-mode variations
# Custom output formats
python sunburst_script.py data.csv --output chart.svg --width 20 --height 20
# Aggregate small slices and control labeling
python sunburst_script.py data.csv --threshold 10.0 --label-threshold 6.0
# Disable adaptive font sizing (use legacy fixed sizes)
python sunburst_script.py data.csv --no-adaptive-font
Key Features:
- Multi-level hierarchy support (3-5 levels)
- Smart color inheritance with two modes (variations/same)
- Top-N aggregation: keep only top N items per level
- Percentage threshold aggregation (local or global mode)
- Flexible label formatting (name-only, name-count, name-percent, full)
- Adaptive font sizing based on segment geometry
- Customizable line widths and label thresholds
- Multiple output formats (PNG, SVG, PDF, EPS, TIFF)
- Dual counting modes (all records or unique values)
Output:
By default generates three files:
- sunburst_chart.png (or specified format)
- sunburst_chart.svg (vector, editable)
- sunburst_chart.pdf (publication-ready)
For detailed documentation, see README.md in this directory or run:
python sunburst_script.py --help
"""
import pandas as pd
import matplotlib.pyplot as plt
import numpy as np
import argparse
import sys
from collections import defaultdict
import matplotlib.patches as mpatches
import matplotlib.colors as mcolors
def load_and_process_data(csv_file, sample_id_col, level_cols, count_unique=False):
"""
Load CSV and process hierarchical data for up to 5 levels
"""
try:
df = pd.read_csv(csv_file, low_memory=False)
print(f"Loaded {len(df)} rows from {csv_file}")
# Filter out None values and empty strings from level_cols
active_levels = [col for col in level_cols if col and col.strip()]
# Filter out rows with missing data in key columns
required_cols = [sample_id_col] + active_levels
for col in required_cols:
if col not in df.columns:
raise ValueError(f"Column '{col}' not found in CSV. Available columns: {list(df.columns)}")
df_clean = df.dropna(subset=required_cols)
print(f"After removing rows with missing data: {len(df_clean)} rows")
print(f"Active hierarchy levels: {len(active_levels)} - {active_levels}")
# Build hierarchical structure dynamically
def build_nested_dict(depth):
if depth == 0:
return set if count_unique else int
return lambda: defaultdict(build_nested_dict(depth - 1))
hierarchy = defaultdict(build_nested_dict(len(active_levels) - 1))
for _, row in df_clean.iterrows():
current_level = hierarchy
for i, level_col in enumerate(active_levels):
key = str(row[level_col]).strip()
if i == len(active_levels) - 1:
if count_unique:
current_level[key].add(str(row[sample_id_col]).strip())
else:
current_level[key] += 1
else:
current_level = current_level[key]
# Convert sets to counts if using count_unique
if count_unique:
def convert_sets_to_counts(obj):
if isinstance(obj, set):
return len(obj)
elif isinstance(obj, defaultdict):
return defaultdict(int, {k: convert_sets_to_counts(v) for k, v in obj.items()})
return obj
hierarchy = convert_sets_to_counts(hierarchy)
# Calculate total unique values
total_count = sum(calculate_total_recursive(v) for v in hierarchy.values())
else:
total_count = len(df_clean)
return hierarchy, total_count, active_levels
except Exception as e:
print(f"Error processing data: {e}")
sys.exit(1)
def aggregate_small_slices(data_dict, threshold_percent, total_for_level, other_label="Other"):
"""
Aggregate small slices into an "Other" category based on percentage threshold
Args:
data_dict: Dictionary of items to potentially aggregate
threshold_percent: Minimum percentage threshold (0-100) for individual items
total_for_level: Total count for the current level
other_label: Label to use for aggregated small items
Returns:
Tuple of (aggregated_dict, other_items_list)
"""
if threshold_percent <= 0:
return data_dict, []
threshold_count = (threshold_percent / 100.0) * total_for_level
# Separate items above and below threshold
main_items = {}
small_items = {}
other_total = 0
for key, value in data_dict.items():
item_count = calculate_total_recursive(value) if not isinstance(value, int) else value
if item_count >= threshold_count:
main_items[key] = value
else:
small_items[key] = value
other_total += item_count
# If we have small items to aggregate and it would be meaningful
if small_items and len(small_items) > 1 and other_total > 0:
# Add the "Other" category
main_items[other_label] = other_total
print(f" Aggregated {len(small_items)} items < {threshold_percent}% into '{other_label}' ({other_total:,} total)")
return main_items, list(small_items.keys())
else:
# Don't aggregate if only one small item or no small items
return data_dict, []
def aggregate_top_n(data_dict, top_n, other_label="Other"):
"""
Keep only top N items by count, aggregate rest into "Other" category
Args:
data_dict: Dictionary of items to potentially aggregate
top_n: Number of top items to keep (None or 0 = no limit)
other_label: Label to use for aggregated small items
Returns:
Tuple of (aggregated_dict, other_items_list)
"""
if not top_n or top_n <= 0:
return data_dict, []
# Calculate totals for sorting
items_with_totals = []
for key, value in data_dict.items():
item_count = calculate_total_recursive(value) if not isinstance(value, int) else value
items_with_totals.append((key, value, item_count))
# Sort by count descending
items_sorted = sorted(items_with_totals, key=lambda x: x[2], reverse=True)
# Split into top N and rest
top_items = items_sorted[:top_n]
rest_items = items_sorted[top_n:]
if not rest_items:
return data_dict, []
# Build result
main_items = {key: value for key, value, _ in top_items}
other_total = sum(count for _, _, count in rest_items)
aggregated_keys = [key for key, _, _ in rest_items]
if other_total > 0:
main_items[other_label] = other_total
print(f" Kept top {top_n} items, aggregated {len(rest_items)} into '{other_label}' ({other_total:,} total)")
return main_items, aggregated_keys
def format_label(key, count, total_samples, level_total, label_style, show_percent):
"""
Format segment label based on style settings
Args:
key: Segment name
count: Segment count
total_samples: Global total for percentage calculation
level_total: Level total for local percentage
label_style: 'name-count' | 'name-only' | 'name-percent' | 'full'
show_percent: Whether to add percentage (can override style)
Returns:
Formatted label string
"""
global_pct = (count / total_samples * 100) if total_samples > 0 else 0
local_pct = (count / level_total * 100) if level_total > 0 else 0
if label_style == 'name-only':
return key
elif label_style == 'name-percent':
return f"{key}\n{global_pct:.1f}%"
elif label_style == 'full':
return f"{key}\n{count:,} ({global_pct:.1f}%)"
elif label_style == 'name-count' or label_style is None:
# Default: name-count, optionally add percent
if show_percent:
return f"{key}\n{count:,} ({global_pct:.1f}%)"
else:
return f"{key}\n{count:,}"
else:
# Fallback
return f"{key}\n{count:,}"
def calculate_adaptive_fontsize(angle_size, ring_width, level, base_min=5, base_max=12):
"""
Calculate font size adaptively based on segment geometry
Args:
angle_size: Angular size of segment in degrees
ring_width: Radial width of the ring
level: Hierarchy level (1-based)
base_min: Minimum font size
base_max: Maximum font size
Returns:
Calculated font size
"""
# Estimate arc length at mid-radius (normalized)
# Larger angles and wider rings = more space = larger font
# Angular factor: scale from 0 (tiny) to 1 (full circle)
angle_factor = min(1.0, angle_size / 90.0) # 90 degrees = max factor
# Ring width factor: thinner rings need smaller text
ring_factor = min(1.0, ring_width / 0.25) # 0.25 is a "normal" ring width
# Level penalty: deeper levels get slightly smaller text
level_penalty = max(0, (level - 1) * 0.5)
# Combined factor
combined = (angle_factor * 0.6 + ring_factor * 0.4)
# Calculate size
fontsize = base_min + (base_max - base_min) * combined - level_penalty
# Clamp to valid range
return max(base_min, min(base_max, fontsize))
def generate_distinct_colors(n_colors):
"""Generate highly distinct colors for better discrimination"""
if n_colors <= 20:
# Use hand-picked distinct colors for small/medium sets
# Order: blues first, then greens, teals, yellows, oranges, reds, purples, greys
base_colors = [
'#2E86AB', # 1 - dark blue
'#45B7D1', # 2 - medium blue
'#85C1E9', # 3 - light blue
'#27AE60', # 4 - green
'#82E0AA', # 5 - light green
'#4ECDC4', # 6 - teal
'#96CEB4', # 7 - sage
'#F4D03F', # 8 - yellow
'#F7DC6F', # 9 - light yellow
'#F8C471', # 10 - orange
'#E74C3C', # 11 - red
'#FF6B6B', # 12 - coral
'#9B59B6', # 13 - purple
'#BB8FCE', # 14 - light purple
'#DDA0DD', # 15 - plum
'#5D6D7E', # 16 - slate grey
'#ABB2B9', # 17 - light grey
'#1ABC9C', # 18 - turquoise
'#E59866', # 19 - tan
'#CD6155', # 20 - dusty red
]
return [base_colors[i % len(base_colors)] for i in range(n_colors)]
else:
# Use multiple colormap cycles for larger sets
colors = []
colormaps = [plt.cm.Set1, plt.cm.Set2, plt.cm.Set3, plt.cm.Paired, plt.cm.Dark2]
for i in range(n_colors):
cmap = colormaps[i // 9 % len(colormaps)]
colors.append(cmap((i % 9) / 9))
return colors
def generate_color_variations(base_color, n_variations):
"""Generate variations of a base color by adjusting brightness"""
if isinstance(base_color, str):
# Convert hex to RGB
if base_color.startswith('#'):
r, g, b = int(base_color[1:3], 16)/255, int(base_color[3:5], 16)/255, int(base_color[5:7], 16)/255
else:
# Assume it's a named color
r, g, b = mcolors.to_rgb(base_color)
else:
# Already RGB tuple
r, g, b = base_color[:3]
variations = []
if n_variations == 1:
return [base_color]
# Generate variations by adjusting brightness
for i in range(n_variations):
# Create variations from 0.3 to 1.0 brightness multiplier
factor = 0.3 + (0.7 * i / (n_variations - 1))
new_r = min(1.0, r * factor + (1 - factor) * 0.9) # Lighten for better contrast
new_g = min(1.0, g * factor + (1 - factor) * 0.9)
new_b = min(1.0, b * factor + (1 - factor) * 0.9)
variations.append((new_r, new_g, new_b))
return variations
def calculate_total_recursive(data):
"""Recursively calculate total from nested dictionary"""
if isinstance(data, int):
return data
return sum(calculate_total_recursive(v) for v in data.values())
def create_sunburst_chart(hierarchy, total_samples, active_levels, output_file='sunburst_chart.png',
title=None, figsize=(18, 18), auto_formats=True,
color_inherit_level=1, color_mode='variations', count_unique=False,
line_width=0.5, threshold_percent=0.0, other_label="Other",
label_threshold=5.0, top_n=None, threshold_mode='local',
label_style='name-count', show_percent=False, adaptive_font=True):
"""
Create a sunburst chart with hierarchical segments for up to 5 levels
Args:
line_width: Width of lines between segments (default: 0.5, original was 2)
threshold_percent: Percentage threshold for aggregating small slices (0 = no aggregation)
other_label: Label to use for aggregated small items
label_threshold: Minimum angle in degrees for showing labels (default: 5.0)
color_inherit_level: Level from which colors should be inherited (1-based indexing)
color_mode: 'variations' = create color variations for deeper levels
'same' = use exact same colors for all levels
top_n: Keep only top N items per level, aggregate rest (None = no limit)
threshold_mode: 'local' = percentage of level total, 'global' = percentage of grand total
label_style: 'name-count' | 'name-only' | 'name-percent' | 'full'
show_percent: Add percentage to labels (simpler toggle, works with name-count)
adaptive_font: Use adaptive font sizing based on segment geometry
"""
fig, ax = plt.subplots(figsize=figsize, subplot_kw=dict(aspect="equal"))
n_levels = len(active_levels)
# Define ring radii dynamically based on number of levels
center_radius = 0.15
ring_width = (0.85 - center_radius) / n_levels
radii = [center_radius + i * ring_width for i in range(n_levels + 1)]
print(f"Creating {n_levels} level sunburst with radii: {radii}")
print(f"Color inheritance level: {color_inherit_level}")
print(f"Color mode: {color_mode}")
print(f"Line width: {line_width}")
print(f"Small slice threshold: {threshold_percent}% ({threshold_mode} mode)")
print(f"Top-N limit: {top_n if top_n else 'None (show all)'}")
print(f"Label threshold: {label_threshold} degrees")
print(f"Label style: {label_style}, show_percent: {show_percent}")
print(f"Adaptive font sizing: {adaptive_font}")
# Apply small slice aggregation to level 1 if threshold is set
processed_hierarchy = dict(hierarchy)
aggregation_info = {} # Track what was aggregated
# First apply top-N if specified
if top_n and top_n > 0:
print(f"Applying top-{top_n} aggregation:")
processed_hierarchy, aggregated_items = aggregate_top_n(
processed_hierarchy, top_n, other_label
)
if aggregated_items:
aggregation_info['level_1_topn'] = aggregated_items
# Then apply percentage threshold if specified
if threshold_percent > 0:
print(f"Applying {threshold_percent}% threshold ({threshold_mode} mode) for small slice aggregation:")
level1_totals = {k: calculate_total_recursive(v) for k, v in processed_hierarchy.items()}
# Use global or local total based on mode
if threshold_mode == 'global':
total_for_threshold = total_samples
else:
total_for_threshold = sum(level1_totals.values())
processed_hierarchy, aggregated_items = aggregate_small_slices(
processed_hierarchy, threshold_percent, total_for_threshold, other_label
)
if aggregated_items:
aggregation_info['level_1'] = aggregated_items
# Calculate totals for level 1 and sort
level1_totals = {k: calculate_total_recursive(v) for k, v in processed_hierarchy.items()}
level1_sorted = sorted(level1_totals.items(), key=lambda x: x[1], reverse=True)
# Generate distinct colors for level 1
level1_colors = generate_distinct_colors(len(level1_sorted))
base_color_map = {key: level1_colors[i] for i, (key, _) in enumerate(level1_sorted)}
segments = [] # Store all segments for drawing
def process_level(data_dict, level, parent_angle_start, parent_angle_size, parent_color, path=[]):
"""Recursively process each level of the hierarchy"""
if level >= n_levels:
return
# Apply aggregation if enabled and we're not at the top level or processing aggregated data
processed_data = data_dict
if level > 0 and other_label not in str(path):
# Calculate total for this level
if isinstance(list(data_dict.values())[0], int):
level_total = sum(data_dict.values())
else:
level_total = sum(calculate_total_recursive(v) for v in data_dict.values())
# Apply top-N first if specified
if top_n and top_n > 0:
processed_data, aggregated_items = aggregate_top_n(
processed_data, top_n, other_label
)
if aggregated_items:
level_key = f"level_{level + 1}_topn"
if level_key not in aggregation_info:
aggregation_info[level_key] = []
aggregation_info[level_key].extend([f"{'/'.join(path)}/{item}" for item in aggregated_items])
# Then apply percentage threshold if specified
if threshold_percent > 0:
# Use global or local total based on mode
if threshold_mode == 'global':
total_for_threshold = total_samples
else:
# Recalculate after top-N aggregation
if isinstance(list(processed_data.values())[0], int):
total_for_threshold = sum(processed_data.values())
else:
total_for_threshold = sum(calculate_total_recursive(v) for v in processed_data.values())
processed_data, aggregated_items = aggregate_small_slices(
processed_data, threshold_percent, total_for_threshold, other_label
)
if aggregated_items:
level_key = f"level_{level + 1}"
if level_key not in aggregation_info:
aggregation_info[level_key] = []
aggregation_info[level_key].extend([f"{'/'.join(path)}/{item}" for item in aggregated_items])
# Sort items by size
if isinstance(list(processed_data.values())[0], int):
# Final level - values are integers
items_sorted = sorted(processed_data.items(), key=lambda x: x[1], reverse=True)
level_total = sum(processed_data.values())
else:
# Intermediate level - values are dictionaries
items_sorted = sorted(processed_data.items(), key=lambda x: calculate_total_recursive(x[1]), reverse=True)
level_total = sum(calculate_total_recursive(v) for v in processed_data.values())
current_angle = parent_angle_start
# Determine color scheme based on inheritance level and mode
level_colors = []
if level + 1 <= color_inherit_level:
# Before or at inheritance level - use distinct colors
if level == 0:
# Level 1 uses predefined distinct colors
level_colors = [base_color_map.get(key, '#CCCCCC') for key, _ in items_sorted]
else:
# Generate distinct colors for this level
level_colors = generate_distinct_colors(len(items_sorted))
else:
# After inheritance level - inherit from parent
if color_mode == 'same':
# Use exact same color as parent
level_colors = [parent_color] * len(items_sorted)
else: # color_mode == 'variations'
# Generate variations of the parent color
level_colors = generate_color_variations(parent_color, len(items_sorted))
# Calculate ring width for this level (for adaptive font sizing)
current_ring_width = radii[level + 1] - radii[level]
for i, (key, value) in enumerate(items_sorted):
if isinstance(value, int):
item_total = value
else:
item_total = calculate_total_recursive(value)
# Calculate angle for this segment
angle_size = (item_total / level_total) * parent_angle_size
# Use assigned color
segment_color = level_colors[i]
# Format label according to style settings
formatted_label = format_label(key, item_total, total_samples, level_total,
label_style, show_percent)
segments.append({
'level': level + 1,
'start_angle': current_angle,
'end_angle': current_angle + angle_size,
'inner_radius': radii[level],
'outer_radius': radii[level + 1],
'color': segment_color,
'label': formatted_label,
'key': key,
'value': item_total,
'path': path + [key],
'ring_width': current_ring_width,
'level_total': level_total
})
# Recursively process next level if it exists
if not isinstance(value, int) and level + 1 < n_levels:
process_level(value, level + 1, current_angle, angle_size, segment_color, path + [key])
current_angle += angle_size
# Start processing from level 1
process_level(processed_hierarchy, 0, 0, 360, None, [])
# Draw all segments
for segment in segments:
# Create wedge with customizable line width
wedge = mpatches.Wedge(
(0, 0), segment['outer_radius'],
segment['start_angle'], segment['end_angle'],
width=segment['outer_radius'] - segment['inner_radius'],
facecolor=segment['color'],
edgecolor='white',
linewidth=line_width # Now customizable
)
ax.add_patch(wedge)
# Calculate if we should show label based on angle size threshold
angle_size = segment['end_angle'] - segment['start_angle']
show_label = angle_size > label_threshold # Apply consistent threshold to all levels
if show_label:
# Calculate label position
mid_angle = (segment['start_angle'] + segment['end_angle']) / 2
mid_radius = (segment['inner_radius'] + segment['outer_radius']) / 2
# Convert to radians
angle_rad = np.radians(mid_angle)
x = mid_radius * np.cos(angle_rad)
y = mid_radius * np.sin(angle_rad)
# Improved text rotation - always radial outward
rotation = mid_angle
# Adjust for readability - text should read outward from center
if mid_angle > 90 and mid_angle <= 270:
rotation = mid_angle + 180 # Flip text in bottom half
# Font size calculation - adaptive or legacy
if adaptive_font:
ring_width = segment.get('ring_width', radii[1] - radii[0])
fontsize = calculate_adaptive_fontsize(angle_size, ring_width, segment['level'])
else:
# Legacy font size calculation
base_fontsize = max(6, min(12, 14 - segment['level']))
if angle_size < 10:
fontsize = max(6, base_fontsize - 2)
else:
fontsize = base_fontsize
fontweight = 'bold' if segment['level'] <= 2 else 'normal'
# Always use black text as requested
text_color = 'black'
# Add text
text = ax.text(x, y, segment['label'],
horizontalalignment='center', verticalalignment='center',
fontsize=fontsize, weight=fontweight, rotation=rotation,
color=text_color)
# Add center circle with total
center_circle = plt.Circle((0, 0), center_radius, fc='white', ec='black', linewidth=3)
ax.add_patch(center_circle)
count_label = 'Unique\nValues' if count_unique else 'Total\nSamples'
ax.text(0, 0, f'{count_label}\n{total_samples:,}',
horizontalalignment='center', verticalalignment='center',
fontsize=14, weight='bold', color='black')
# Set equal aspect ratio and remove axes
max_radius = radii[-1]
ax.set_xlim(-max_radius * 1.1, max_radius * 1.1)
ax.set_ylim(-max_radius * 1.1, max_radius * 1.1)
ax.set_aspect('equal')
ax.axis('off')
# Set title only if provided
if title:
plt.title(title, fontsize=18, weight='bold', pad=20, color='black')
# Save the figure in specified format(s)
plt.tight_layout()
# Determine output format from file extension
file_ext = output_file.lower().split('.')[-1]
# Set appropriate DPI and format parameters
save_params = {
'bbox_inches': 'tight',
'facecolor': 'white'
}
if file_ext in ['png', 'jpg', 'jpeg', 'tiff', 'tif']:
save_params['dpi'] = 300
elif file_ext in ['svg', 'pdf', 'eps']:
# Vector formats don't need DPI but benefit from other settings
save_params['dpi'] = 300 # Still good for any embedded raster elements
if file_ext == 'svg':
save_params['format'] = 'svg'
elif file_ext == 'eps':
save_params['format'] = 'eps'
plt.savefig(output_file, **save_params)
print(f"Sunburst chart saved as {output_file} ({file_ext.upper()} format)")
# Also save in additional formats if specified
if auto_formats:
base_name = '.'.join(output_file.split('.')[:-1])
# Automatically save SVG version for editing (unless already SVG)
if file_ext != 'svg':
svg_file = f"{base_name}.svg"
plt.savefig(svg_file, format='svg', bbox_inches='tight', facecolor='white')
print(f"Also saved editable SVG version: {svg_file}")
# Automatically save PDF version for high-quality printing (unless already PDF)
if file_ext != 'pdf':
pdf_file = f"{base_name}.pdf"
plt.savefig(pdf_file, format='pdf', bbox_inches='tight', facecolor='white')
print(f"Also saved PDF version: {pdf_file}")
# Display summary statistics
print(f"\nSummary Statistics:")
count_type = "unique values" if count_unique else "samples"
print(f"Total {count_type}: {total_samples:,}")
print(f"Number of levels: {n_levels}")
print(f"Color inheritance from level: {color_inherit_level}")
print(f"Color mode: {color_mode}")
print(f"Line width: {line_width}")
print(f"Label threshold: {label_threshold} degrees")
print(f"Label style: {label_style}")
print(f"Top-N: {top_n if top_n else 'disabled'}")
print(f"Threshold mode: {threshold_mode}")
print(f"Adaptive font: {adaptive_font}")
# Report aggregation results
if aggregation_info:
print(f"\nSmall slice aggregation (threshold: {threshold_percent}%):")
for level_key, items in aggregation_info.items():
level_num = level_key.split('_')[1]
print(f" Level {level_num}: {len(items)} items aggregated into '{other_label}'")
for item in items[:5]: # Show first 5 items
print(f" - {item}")
if len(items) > 5:
print(f" ... and {len(items) - 5} more")
for i, level_name in enumerate(active_levels):
level_segments = [s for s in segments if s['level'] == i + 1]
print(f"Level {i+1} ({level_name}): {len(level_segments)} categories")
for key, total in level1_sorted:
percentage = (total / total_samples) * 100
print(f" {key}: {total:,} {count_type} ({percentage:.1f}%)")
def main():
parser = argparse.ArgumentParser(description='Generate sunburst chart from CSV data (up to 5 levels)')
parser.add_argument('csv_file', help='Path to input CSV file')
parser.add_argument('--sample-id', default='Sample-ID', help='Column name for sample IDs (default: Sample-ID)')
parser.add_argument('--level1', default='Partner_sub', help='Column for level 1 (default: Partner_sub)')
parser.add_argument('--level2', default='partner', help='Column for level 2 (default: partner)')
parser.add_argument('--level3', default='Project-Code', help='Column for level 3 (default: Project-Code)')
parser.add_argument('--level4', default=None, help='Column for level 4 (optional)')
parser.add_argument('--level5', default=None, help='Column for level 5 (optional)')
parser.add_argument('--color-inherit-level', type=int, default=1,
help='Level from which colors should be inherited (1-5). ' +
'Level 1: each top-level category and descendants get unique colors. ' +
'Level 2: levels 1-2 get unique colors, level 3+ inherit from level 2, etc. (default: 1)')
parser.add_argument('--color-mode', choices=['variations', 'same'], default='variations',
help='Color inheritance mode: "variations" creates color shades for deeper levels, ' +
'"same" uses identical colors for all inherited levels (default: variations)')
parser.add_argument('--count-unique', action='store_true',
help='Count unique values in sample-id column instead of all records (default: False)')
parser.add_argument('--output', default='sunburst_chart.png',
help='Output filename with extension (default: sunburst_chart.png)\n' +
'Supported formats: PNG, JPG, PDF, SVG, EPS, TIFF\n' +
'SVG and PDF are automatically generated for editing')
parser.add_argument('--title', default=None, help='Chart title (default: no title)')
parser.add_argument('--width', type=int, default=18, help='Figure width in inches (default: 18)')
parser.add_argument('--height', type=int, default=18, help='Figure height in inches (default: 18)')
parser.add_argument('--no-auto-formats', action='store_true',
help='Skip automatic generation of SVG and PDF versions')
# NEW ARGUMENTS for enhancements
parser.add_argument('--line-width', type=float, default=0.5,
help='Width of lines between segments (default: 0.5, original was 2.0)')
parser.add_argument('--threshold', type=float, default=0.0,
help='Percentage threshold for aggregating small slices into "Other" group (0-100, default: 0 = no aggregation)')
parser.add_argument('--other-label', default='Other',
help='Label to use for aggregated small items (default: "Other")')
parser.add_argument('--label-threshold', type=float, default=5.0,
help='Minimum angle in degrees for showing segment labels (default: 5.0)')
# NEW: Top-N, threshold-mode, label-style, show-percent, adaptive-font
parser.add_argument('--top-n', type=int, default=None,
help='Keep only top N items per level, aggregate rest into "Other" (default: None = show all)')
parser.add_argument('--threshold-mode', choices=['local', 'global'], default='local',
help='Threshold calculation mode: "local" = percentage of level total, '
'"global" = percentage of grand total (default: local)')
parser.add_argument('--label-style', choices=['name-count', 'name-only', 'name-percent', 'full'],
default='name-count',
help='Label format style: "name-count" (default), "name-only", "name-percent", "full" (name + count + percent)')
parser.add_argument('--show-percent', action='store_true',
help='Add percentage to labels (works with name-count style)')
parser.add_argument('--no-adaptive-font', action='store_true',
help='Disable adaptive font sizing (use legacy fixed sizing)')
args = parser.parse_args()
# Validate threshold
if args.threshold < 0 or args.threshold > 100:
print(f"Error: --threshold must be between 0 and 100 (got {args.threshold})")
sys.exit(1)
# Validate color inheritance level
level_cols = [args.level1, args.level2, args.level3, args.level4, args.level5]
active_levels = [col for col in level_cols if col and col.strip()]
if args.color_inherit_level < 1 or args.color_inherit_level > len(active_levels):
print(f"Error: --color-inherit-level must be between 1 and {len(active_levels)} (number of active levels)")
sys.exit(1)
# Load and process data
hierarchy, total_samples, active_levels = load_and_process_data(
args.csv_file, args.sample_id, level_cols, args.count_unique
)
# Create the chart
create_sunburst_chart(
hierarchy, total_samples, active_levels,
output_file=args.output,
title=args.title,
figsize=(args.width, args.height),
auto_formats=not args.no_auto_formats,
color_inherit_level=args.color_inherit_level,
color_mode=args.color_mode,
count_unique=args.count_unique,
line_width=args.line_width,
threshold_percent=args.threshold,
other_label=args.other_label,
label_threshold=args.label_threshold,
top_n=args.top_n,
threshold_mode=args.threshold_mode,
label_style=args.label_style,
show_percent=args.show_percent,
adaptive_font=not args.no_adaptive_font
)
if __name__ == "__main__":
# Example usage if run directly
if len(sys.argv) == 1:
print("Enhanced Sunburst Chart Generator")
print("=================================")
print("\nBasic usage:")
print(" python sunburst_script.py bge_museum_data.csv")
print("\nGrouping options:")
print(" python sunburst_script.py data.csv --top-n 8 # Keep top 8, group rest")
print(" python sunburst_script.py data.csv --threshold 5.0 # Group items < 5%")
print(" python sunburst_script.py data.csv --threshold-mode global # Use global totals")
print(" python sunburst_script.py data.csv --top-n 10 --threshold 2 # Combined")
print("\nLabel formatting:")
print(" python sunburst_script.py data.csv --label-style name-only # Names only")
print(" python sunburst_script.py data.csv --label-style name-percent # Names + %")
print(" python sunburst_script.py data.csv --label-style full # Name + count + %")
print(" python sunburst_script.py data.csv --show-percent # Add % to default")
print("\nVisual options:")
print(" python sunburst_script.py data.csv --line-width 0.2 # Thinner lines")
print(" python sunburst_script.py data.csv --label-threshold 3.0 # More labels visible")
print(" python sunburst_script.py data.csv --no-adaptive-font # Legacy font sizing")
print("\nCombined example:")
print(" python sunburst_script.py data.csv --top-n 10 --label-style full --line-width 0.3")
print("\nAll original features still supported:")
print(" python sunburst_script.py data.csv --output chart.svg")
print(" python sunburst_script.py data.csv --level4 Category4 --level5 Category5")
print(" python sunburst_script.py data.csv --color-inherit-level 2")
print(" python sunburst_script.py data.csv --count-unique")
print("\nSupported formats: PNG, JPG, PDF, SVG, EPS, TIFF")
print("Note: SVG and PDF versions are automatically created for editing")
print("\nFor help: python sunburst_script.py --help")
else:
main()