File size: 8,590 Bytes
a10ba7f | 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 | import matplotlib.pyplot as plt
import matplotlib.patches as mpatches
import matplotlib.lines as mlines
import numpy as np
from adjustText import adjust_text
# ββ 1. DATA βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
data = [
dict(name='Mnist\n(2015)', mae=7.29, gflops=0.10, params=1.82, group='prior'),
dict(name='iTracker\n(2016)', mae=7.67, gflops=3.97, params=6.28, group='prior'),
dict(name='GazeNet\n(2017)', mae=6.62, gflops=72.24, params=90.23, group='prior'),
dict(name='FullFace\n(2017)', mae=5.65, gflops=29.90, params=190.00, group='prior'),
dict(name='RT-Gene\n(2018)', mae=5.36, gflops=12.21, params=31.66, group='prior'),
dict(name='DilatedNet\n(2019)', mae=5.07, gflops=202.00, params=3.92, group='prior'),
dict(name='Gaze360\n(2019)', mae=4.66, gflops=3.65, params=11.72, group='recent'),
dict(name='FAR-Net\n(2021)', mae=5.12, gflops=0.65, params=1.94, group='recent'),
dict(name='FR-Net\n(2024)', mae=4.95, gflops=0.15, params=0.85, group='recent'),
dict(name='FGI-Net\n(2025)', mae=4.81, gflops=0.08, params=0.45, group='recent'),
dict(name='Heavy Teacher\n(ResNet50)', mae=4.15, gflops=4.12, params=25.56, group='teacher'),
dict(name='LIPE \n(Ours)', mae=4.72, gflops=0.02125, params=0.18, group='ours'),
]
# ββ 2. STYLE ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
plt.rcParams['font.family'] = 'DejaVu Sans'
plt.rcParams['xtick.direction'] = 'in'
plt.rcParams['ytick.direction'] = 'in'
GROUP_STYLE = {
'prior': {'color': '#C8C5BC', 'edgecolor': '#5F5E5A', 'label': 'Prior Works (2015β2019)'},
'recent': {'color': '#6B8E23', 'edgecolor': '#3A5010', 'label': 'Recent Edge SOTA'},
'teacher': {'color': '#9E9E9E', 'edgecolor': '#555555', 'label': 'Heavy Teacher Baseline'},
'ours': {'color': '#E53935', 'edgecolor': '#7B1FA2', 'label': 'LIPE (Ours)'},
}
ARROW_STYLE = dict(
arrowstyle='->',
color='#888888',
lw=0.9,
connectionstyle='arc3,rad=0.0'
)
fig, ax = plt.subplots(figsize=(9, 6.5), dpi=200)
fig.patch.set_facecolor('#FAFAFA')
ax.set_facecolor('#FAFAFA')
def bubble_area(p):
return 90 + 380 * np.log10(p + 1)
# ββ 3. SCATTER POINTS βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
for d in data:
style = GROUP_STYLE[d['group']]
s = bubble_area(d['params'])
marker = 'D' if d['group'] == 'ours' else 'o'
lw = 2.0 if d['group'] == 'ours' else 0.8
zo = 6 if d['group'] == 'ours' else (5 if d['group'] == 'teacher' else 3)
ax.scatter(
d['gflops'], d['mae'],
s=s,
color=style['color'],
edgecolors=style['edgecolor'],
marker=marker,
alpha=0.88,
linewidths=lw,
zorder=zo,
)
# ββ 4. ANNOTATIONS WITH ARROWS (adjustText for collision avoidance) βββββββββββ
# Manual offsets (in data-space offsets via display transform) β fine-tuned per label
# key: point name (first line), value: (dx_pts, dy_pts) offset for text
MANUAL_OFFSETS = {
'LIPE \n(Ours)': (-68, 32),
'FGI-Net\n(2025)': ( 12, -28),
'FR-Net\n(2024)': ( 12, 22),
'FAR-Net\n(2021)': (-70, -12),
'Gaze360\n(2019)': ( 14, 22),
'Heavy Teacher\n(ResNet50)': ( 14, 10),
'DilatedNet\n(2019)': ( 14, 0),
'RT-Gene\n(2018)': ( 14, 10),
'FullFace\n(2017)': ( 14, -26),
'GazeNet\n(2017)': ( 14, 0),
'iTracker\n(2016)': (-72, 10),
'Mnist\n(2015)': ( 14, -26),
}
texts = []
arrows = []
for d in data:
is_ours = (d['group'] == 'ours')
fw = 'bold' if is_ours else 'normal'
col = '#C62828' if is_ours else '#333333'
fs = 8.5 if is_ours else 7.5
bg_alpha = 0.82 if is_ours else 0.70
bg_color = '#FFF9F9' if is_ours else 'white'
dx, dy = MANUAL_OFFSETS.get(d['name'], (12, 0))
ann = ax.annotate(
d['name'],
xy=(d['gflops'], d['mae']),
xytext=(dx, dy),
textcoords='offset points',
fontsize=fs,
fontweight=fw,
color=col,
ha='center',
va='center',
zorder=9,
bbox=dict(
boxstyle='round,pad=0.28',
fc=bg_color,
ec='#CCCCCC' if not is_ours else '#E57373',
lw=0.6 if not is_ours else 1.0,
alpha=bg_alpha,
),
arrowprops=dict(
arrowstyle='->',
color='#AAAAAA' if not is_ours else '#E53935',
lw=0.85 if not is_ours else 1.2,
connectionstyle='arc3,rad=0.15',
),
)
texts.append(ann)
# ββ 5. PARETO FRONTIER ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
pareto_x = [0.02125, 0.08, 3.65, 4.12]
pareto_y = [4.72, 4.81, 4.66, 4.15]
ax.plot(pareto_x, pareto_y, color='#888888', linestyle='--', linewidth=1.3, zorder=2, alpha=0.8)
# ββ 6. AXES βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
ax.set_xscale('log')
ax.set_xlabel('Computational Complexity (GFLOPs) [Log Scale]',
fontsize=10, fontweight='bold', labelpad=7, color='#2E2E2E')
ax.set_ylabel('Gaze Estimation Error (MAE in Degrees) [Lower is Better]',
fontsize=10, fontweight='bold', labelpad=7, color='#2E2E2E')
ax.set_xlim(0.005, 600.0)
ax.set_ylim(8.2, 3.5)
ax.grid(True, which='both', ls=':', color='#DDDDDD', zorder=1)
ax.tick_params(axis='both', colors='#555555', labelsize=8.5)
for spine in ax.spines.values():
spine.set_edgecolor('#CCCCCC')
spine.set_linewidth(0.7)
# ββ 7. LEGENDS ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
# Left legend: Model classification
group_handles = []
for g in ['ours', 'recent', 'teacher', 'prior']:
st = GROUP_STYLE[g]
mk = 'D' if g == 'ours' else 'o'
group_handles.append(mlines.Line2D(
[], [], color='none', marker=mk, markersize=7,
markerfacecolor=st['color'], markeredgecolor=st['edgecolor'],
markeredgewidth=1.1, label=st['label'],
))
group_handles.append(mlines.Line2D(
[], [], color='#888888', linestyle='--', linewidth=1.3,
label='Current Pareto Frontier',
))
leg1 = ax.legend(
handles=group_handles,
loc='lower left',
title='Model Classification',
fontsize=8, title_fontsize=8.5,
frameon=True, facecolor='white', edgecolor='#CCCCCC', framealpha=0.93,
borderpad=0.7, labelspacing=0.5,
bbox_to_anchor=(0.01, 0.01),
)
ax.add_artist(leg1)
# Right legend: Bubble size
sizes_demo = [0.18, 1.94, 11.72, 190.0]
size_labels = ['0.18M (Ours)', '1.94M', '11.72M', '190.0M']
size_handles = [
plt.scatter([], [],
s=bubble_area(sz) * 0.50,
color='#C8C5BC', alpha=0.7,
edgecolors='#5F5E5A', marker='o')
for sz in sizes_demo
]
leg2 = ax.legend(
handles=size_handles, labels=size_labels,
loc='lower right',
title='Bubble Size (# Params)',
fontsize=8, title_fontsize=8.5,
frameon=True, facecolor='white', edgecolor='#CCCCCC', framealpha=0.93,
labelspacing=1.2, borderpad=0.9, handletextpad=1.1,
bbox_to_anchor=(0.99, 0.01),
)
# ββ 8. TITLE ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
ax.set_title(
'Accuracy vs. Efficiency Trade-off: Gaze Estimation Benchmark',
fontsize=11, fontweight='bold', color='#1A1A1A', pad=10,
)
plt.tight_layout(pad=1.4)
plt.savefig('/mnt/user-data/outputs/Figure_1_v2.png', dpi=300, bbox_inches='tight',
facecolor=fig.get_facecolor())
print("Saved Figure_1_v2.png")
plt.show()
|