Перейти к содержанию

Многоклассовый классификатор

Модуль cnn_neat/multiclass_head.py даёт многоклассовый считывание (readout) для геномов с несколькими выходными «наконечниками» (output_nodes), которые появляются после слияния OVA-чемпионов (OVA-merge). Он включается при ExperimentConfig.task_mode="multiclass" и заменяет один BCE-логит бинарной головы на \(K\) логитов и cross-entropy.

Логиты из tip'ов

Каждый выходной конец \(k\) даёт карту признаков \(\mathbf{F}^{(k)}_b \in \mathbb{R}^{C\times H\times W}\). Тем же правилом, что и бинарная голова (output_to_score, режим mean или mean_minus_median), карта сворачивается в скаляр \(s_k\). Скаляры складываются в вектор логитов:

\[ \mathbf{z}_b = \big(s_0, s_1, \dots, s_{K-1}\big) \in \mathbb{R}^{K}, \qquad s_k = \mathrm{readout}\big(\mathbf{F}^{(k)}_b\big). \]

Ключевые функции:

Функция Роль
multiclass_logits_from_outputs(feats, output_readout=...) Стек скаляров tip'ов → (B, K)
multiclass_logits(genome, images, output_readout=...) forward_all_outputs + стек
multiclass_classification_loss(genome, images, targets, ...) Cross-entropy по логитам tip'ов
predict_multiclass_from_logits(logits) argmax по классам
compute_multiclass_metrics(y_true, y_pred, num_classes) accuracy + macro balanced accuracy

Предсказание и обучение

  • Предсказание — winner-take-all по tip'ам (без обучаемого softmax-слияния):
\[ \hat{y}_b = \arg\max_{k} \; s_k . \]
  • Обучение (periodic GD) — стандартная многоклассовая кросс-энтропия по логитам:
\[ \mathcal{L} = \frac{1}{B}\sum_b -\log \frac{e^{z_{b,y_b}}}{\sum_{k} e^{z_{b,k}}} . \]
  • Приспособленность — macro balanced accuracy (среднее по классам recall):
\[ \mathrm{BalAcc} = \frac{1}{K}\sum_{k=0}^{K-1} \frac{\#\{b:\, y_b=k \wedge \hat{y}_b=k\}}{\#\{b:\, y_b=k\}} . \]

compute_multiclass_metrics возвращает accuracy, balanced_accuracy и fitness (= balanced accuracy).

Считывание по датасету

Для CIFAR-10 используется mean_minus_median, для MNIST — mean: при одноканальном 1×1 output разность «mean − median» вырождается в 0, и приспособленность (fitness) залипает на 0.5 (см. MNIST).

Интеграция: приспособленность — в eval.py (_multiclass_tip_metrics, _evaluate_fitness_multiclass_gpu), GD-лосс — в periodic_gd.py (ветка task_mode=="multiclass"). Требует network.output_nodes и forward_all_outputs() (multi-tip DAG из OVA-merge).