Многоклассовый классификатор¶
Модуль 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\). Скаляры складываются в
вектор логитов:
Ключевые функции:
| Функция | Роль |
|---|---|
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-слияния):
- Обучение (periodic GD) — стандартная многоклассовая кросс-энтропия по логитам:
- Приспособленность — macro balanced accuracy (среднее по классам recall):
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).