#!/usr/bin/env python """SpiceFusionNet architecture figure, TResNet-food reference style: 3D tensor slabs, small pastel boxes, an MBConv-block call-out above the flow, a snaking backbone so the three streams converge cleanly, operation circles, and class-circle outputs. Saves outputs/spicefusionnet_arch.{png,pdf}. """ from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt from matplotlib.patches import FancyBboxPatch, FancyArrowPatch, Rectangle, Circle, Polygon from matplotlib.colors import to_rgb matplotlib.rcParams.update({"font.family": "DejaVu Sans", "font.size": 7}) ROOT = Path("/mnt/d/SpiceNet") if Path("/mnt/d/SpiceNet").exists() else Path("D:/SpiceNet") BLOCK, PINK, YEL, GRN, PUR = "#c9dcf2", "#f2c9c9", "#f6e6b0", "#d3ecd0", "#dccdee" SLAB, EDGE, ARR = "#e9eef5", "#555", "#555" def _d(c, f): r, g, b = to_rgb(c); return (r * f, g * f, b * f) fig, ax = plt.subplots(figsize=(13, 6.7)) ax.set_xlim(0, 13); ax.set_ylim(2.1, 8.9); ax.axis("off") def box(x, y, w, h, label, fc=BLOCK, fs=6.6): ax.add_patch(FancyBboxPatch((x, y), w, h, boxstyle="round,pad=0.006,rounding_size=0.04", fc=fc, ec=EDGE, lw=0.8, mutation_scale=1, zorder=4)) ax.text(x + w / 2, y + h / 2, label, ha="center", va="center", fontsize=fs, zorder=5) def slab(x, y, w, h, fc=SLAB, d=0.15, label="", n=1, fs=6.6): for i in range(n - 1, -1, -1): ox, oy = x + i * 0.1, y + i * 0.1 ax.add_patch(Polygon([(ox, oy + h), (ox + w, oy + h), (ox + w + d, oy + h + d), (ox + d, oy + h + d)], fc=_d(fc, .9), ec=EDGE, lw=0.7, zorder=3 + n - i)) ax.add_patch(Polygon([(ox + w, oy), (ox + w + d, oy + d), (ox + w + d, oy + h + d), (ox + w, oy + h)], fc=_d(fc, .78), ec=EDGE, lw=0.7, zorder=3 + n - i)) ax.add_patch(Rectangle((ox, oy), w, h, fc=fc, ec=EDGE, lw=0.8, zorder=3 + n - i)) if label: ax.text(x + w / 2, y - 0.12, label, ha="center", va="top", fontsize=fs) def opc(x, y, s, r=0.13): ax.add_patch(Circle((x, y), r, fc="white", ec=EDGE, lw=0.9, zorder=6)) ax.text(x, y, s, ha="center", va="center", fontsize=8, zorder=7) def ar(x1, y1, x2, y2, lw=0.85, ls="-"): ax.add_patch(FancyArrowPatch((x1, y1), (x2, y2), arrowstyle="-|>", mutation_scale=7, lw=lw, ls=ls, color=ARR, shrinkA=1.5, shrinkB=1.5, zorder=2)) bw = 1.15 # ---------------- MBConv call-out (top) ---------------- cy = 8.35 mx = 3.6; prev = None for lab, w, fc in [("$1{\\times}1$ exp", 0.9, BLOCK), ("DWConv $k$", 0.9, BLOCK), ("SE", 0.55, YEL), ("$1{\\times}1$ proj", 0.9, BLOCK)]: box(mx, cy - 0.25, w, 0.5, lab, fc, 6.3) if prev: ar(prev, cy, mx, cy) prev = mx + w; mx += w + 0.28 opc(mx + 0.02, cy, "$+$") ax.annotate("", xy=(mx + 0.02, cy + 0.12), xytext=(3.8, cy + 0.25), arrowprops=dict(arrowstyle="-|>", color=ARR, lw=0.8, connectionstyle="arc3,rad=-0.45")) ax.text(3.5, cy, "MBConv", ha="right", va="center", fontsize=6.6, style="italic") # ---------------- backbone, snaking two rows ---------------- y1, y2 = 7.15, 5.55 slab(0.25, y1 - 0.5, 0.5, 1.0, SLAB, label="input\n$224^2$", n=3, fs=6.3) row1 = [("Stem", "$3{\\times}3$/2·48", PINK), ("S1", "MBC1 k3·24$\\times$2", BLOCK), ("S2", "MBC6 k3·32$\\times$4", BLOCK), ("S3", "MBC6 k5·56$\\times$4", BLOCK), ("S4", "MBC6 k3·112$\\times$6", BLOCK)] row2 = [("S5", "MBC6 k5·160$\\times$6", BLOCK), ("S6", "MBC6 k5·272$\\times$8", BLOCK), ("S7", "MBC6 k3·448$\\times$2", BLOCK), ("Head", "$1{\\times}1$+GAP·1792", PINK)] bx = 1.3; prev = 0.9; s2c = None for name, cfg, fc in row1: box(bx, y1 - 0.38, bw, 0.76, f"$\\bf{{{name}}}$\n{cfg}", fc, 6.1) ar(prev, y1, bx, y1) if name == "S2": s2c = bx + bw / 2 prev = bx + bw; bx += bw + 0.05 xr = prev # right edge of row1 ax.annotate("", xy=(xr - bw / 2, y2 + 0.38), xytext=(xr - bw / 2, y1 - 0.38), arrowprops=dict(arrowstyle="-|>", color=ARR, lw=0.85, connectionstyle="arc3,rad=0.0")) # turn down ax.annotate("", xy=(s2c, cy - 0.25), xytext=(s2c, y1 + 0.38), arrowprops=dict(arrowstyle="<->", color=EDGE, lw=0.9)) # callout link bx = xr - bw; prev = None for name, cfg, fc in row2: # right to left box(bx - bw, y2 - 0.38, bw, 0.76, f"$\\bf{{{name}}}$\n{cfg}", fc, 6.1) if prev is not None: ar(prev, y2, bx, y2) prev = bx - bw; bx -= bw + 0.05 ar(prev, y2, prev - 0.28, y2) slab(prev - 0.6, y2 - 0.5, 0.3, 1.0, "#f2c79a", label="$f_{\\mathrm{cnn}}$\n1792", fs=6.3) fcnn_c = (prev - 0.45, y2) # ---------------- descriptor branches (lower left) ---------------- yT, yK = 3.9, 2.85 box(1.3, yT - 0.26, 1.5, 0.52, "LBP+GLCM 58-d", YEL, 6.2) box(2.95, yT - 0.26, 1.05, 0.52, "MLP$\\to$256", YEL, 6.2) slab(4.2, yT - 0.4, 0.3, 0.8, "#f0d9a8", label="$f_{\\mathrm{tex}}$\n256", fs=6.2) box(1.3, yK - 0.26, 1.5, 0.52, "HSV hist 100-d", GRN, 6.2) box(2.95, yK - 0.26, 1.05, 0.52, "MLP$\\to$128", GRN, 6.2) slab(4.2, yK - 0.4, 0.3, 0.8, "#c9e3c6", label="$f_{\\mathrm{col}}$\n128", fs=6.2) ax.text(2.6, 4.55, "hand-built descriptors of the input image", fontsize=6.2, style="italic", ha="center") ar(2.8, yT, 2.95, yT); ar(4.0, yT, 4.2, yT) ar(2.8, yK, 2.95, yK); ar(4.0, yK, 4.2, yK) # ---------------- attention fusion + classifier + output ---------------- yf = 3.35 fx0 = 5.2 ax.add_patch(FancyBboxPatch((fx0, yf - 1.05), 2.75, 2.1, boxstyle="round,pad=0.02", fc="#f4f1f8", ec=EDGE, lw=0.85, ls="--", zorder=1)) ax.text(fx0 + 1.37, yf + 0.88, "attention fusion", ha="center", fontsize=6.6, style="italic") ar(fcnn_c[0], y2 - 0.5, fx0, yf + 0.55) # f_cnn -> fusion ar(4.5, yT, fx0, yf + 0.1) # f_tex -> fusion ar(4.5, yK, fx0, yf - 0.4) # f_col -> fusion box(fx0 + 0.15, yf + 0.3, 1.2, 0.4, "concat 2176", "white", 6.2) box(fx0 + 1.45, yf + 0.3, 1.15, 0.4, "gate softmax", "white", 6.2) ar(fx0 + 1.35, yf + 0.5, fx0 + 1.45, yf + 0.5) opc(fx0 + 2.0, yf - 0.4, "$\\times$") ar(fx0 + 2.0, yf + 0.3, fx0 + 2.0, yf - 0.28) box(fx0 + 0.15, yf - 0.6, 1.2, 0.4, "weighted sum", "white", 6.2) ar(fx0 + 1.87, yf - 0.4, fx0 + 1.35, yf - 0.4) slab(fx0 + 2.85, yf - 0.4, 0.26, 0.8, "#d8c9ea", label="fused\n2176", fs=6.2) ar(fx0 + 2.62, yf, fx0 + 2.85, yf) cx = 8.55 ax.text(cx + 1.0, yf + 0.75, "classifier", ha="center", fontsize=6.6, style="italic") ar(fx0 + 3.15, yf, cx, yf) cbx = cx; prev = None for lab, c in [("FC$\\to$512", PUR), ("BN·ReLU", PUR), ("Drop", PUR), ("FC$\\to$22", PUR), ("soft\nmax", YEL)]: box(cbx, yf - 0.24, 0.82, 0.48, lab, c, 6.0) if prev is not None: ar(prev, yf, cbx, yf) prev = cbx + 0.82; cbx += 0.82 + 0.05 ox = prev + 0.55 ax.text(ox, yf + 0.95, "22 classes", ha="center", fontsize=6.6, style="italic") for k, (name, yo) in enumerate([("coriander", yf + 0.62), ("cumin", yf + 0.2), ("$\\vdots$", yf - 0.18), ("nutmeg", yf - 0.55)]): if name != "$\\vdots$": ax.add_patch(Circle((ox, yo), 0.12, fc="#e8b6b0" if k == 0 else "white", ec=EDGE, lw=0.8, zorder=5)) ar(prev, yf, ox - 0.12, yo, lw=0.7) ax.text(ox + 0.2, yo, name, ha="left", va="center", fontsize=6.3, color="#c0504d" if k == 0 else "#333", fontweight="bold" if k == 0 else "normal") else: ax.text(ox, yo, name, ha="center", va="center", fontsize=8) fig.tight_layout(pad=0.3) out = ROOT / "outputs" / "spicefusionnet_arch" fig.savefig(str(out) + ".png", dpi=300, bbox_inches="tight") fig.savefig(str(out) + ".pdf", bbox_inches="tight") print("saved", out)