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

Pooling merge

3.5 Адаптивный пулинг и выравнивание карт

После свёртки на ребре \(u \to w\) применяется AdaptiveAvgPool2d к целевому \((H_w, W_w)\):

\[ \mathbf{Y}' = \text{AdaptiveAvgPool2d}(\mathbf{Y},\; (H_w, W_w)) \]

PyTorch разбивает \(H' \times W'\) на сетку \(H_w \times W_w\) прямоугольных бинов \(\mathcal{R}_{i,j}\) и вычисляет среднее внутри каждого:

\[ Y'_{b,c,i,j} = \frac{1}{|\mathcal{R}_{i,j}|} \sum_{(h,w) \in \mathcal{R}_{i,j}} Y_{b,c,h,w} \]

Это гарантирует точное совпадение размеров с node_sizes[w] независимо от промежуточного \(H'\).

Примечание для режима edge_weights=True: при включённых весовых коэффициентах рёбер AdaptiveAvgPool2d заменяется на операцию кроп/пад (_align_feature_map) — это связано с необходимостью дифференцировать по скалярным весам рёбер.



3.6 Слияние нескольких входов (merge-узел)

Если в вершину \(w\) входят \(K \geq 1\) активных рёбер от вершин \(u_1, \ldots, u_K\):

\[ \mathbf{A}_w = \frac{1}{K} \sum_{k=1}^{K} \text{Pool}_{(H_w, W_w)}\!\Big(\text{Conv}_{\theta_{u_k, w}}(\mathbf{A}_{u_k})\Big) \]

Реализация в execute_plan:

# K > 1
torch.stack(parts, dim=0).mean(dim=0)
# K == 1
parts[0]

При edge_weights=True — softmax-взвешенная выпуклая комбинация:

\[ \mathbf{A}_w = \sum_{k=1}^{K} \alpha_k \cdot \text{Align}_{(H_w, W_w)}\!\Big(\text{Conv}_{\theta_{u_k,w}}(\mathbf{A}_{u_k})\Big), \qquad \alpha_k = \frac{e^{s_k}}{\sum_{j=1}^K e^{s_j}} \]

где \(s_k\) — обучаемые скалярные параметры ребра (не входят в weight_scale, учатся при fine-tune).

Линейность: merge-узел при \(\alpha_k = 1/K\) — линейный комбинатор.