rslearn.models.attention_pooling¶
attention_pooling ¶
An attention pooling layer.
SimpleAttentionPool ¶
Bases: IntermediateComponent
Simple Attention Pooling.
Given a token feature map of shape BCHWN, learn an attention layer which aggregates over the N dimension.
This is done simply by learning a mapping D->1 which is the weight which should be assigned to each token during averaging:
output = sum [feat_token * W(feat_token) for feat_token in feat_tokens]
Source code in rslearn/models/attention_pooling.py
forward_for_map ¶
Attention pooling for a single feature map (BCHWN tensor).
Source code in rslearn/models/attention_pooling.py
forward ¶
forward(intermediates: Any, context: ModelContext) -> FeatureMaps
Forward pass for attention pooling linear probe.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
intermediates
|
Any
|
the output from the previous component, which must be a TokenFeatureMaps. We pool over the final dimension in the TokenFeatureMaps. If multiple maps are passed, we apply the same linear layers to all of them. |
required |
context
|
ModelContext
|
the model context. |
required |
Returns:
| Type | Description |
|---|---|
FeatureMaps
|
torch.Tensor: - output, attentioned pool over the last dimension (B, C, H, W) |
Source code in rslearn/models/attention_pooling.py
AttentionPool ¶
Bases: IntermediateComponent
Attention Pooling.
Given a feature map of shape BCHWN, learn an attention layer which aggregates over the N dimension.
We do this by learning a query token, and applying a standard attention mechanism against this learned query token.
Source code in rslearn/models/attention_pooling.py
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 | |
init_weights ¶
forward_for_map ¶
Attention pooling for a single feature map (BCHWN tensor).
Source code in rslearn/models/attention_pooling.py
forward ¶
forward(intermediates: Any, context: ModelContext) -> FeatureMaps
Forward pass for attention pooling linear probe.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
intermediates
|
Any
|
the output from the previous component, which must be a TokenFeatureMaps. We pool over the final dimension in the TokenFeatureMaps. If multiple feature maps are passed, we apply the same attention weights (query token and linear k, v layers) to all the maps. |
required |
context
|
ModelContext
|
the model context. |
required |
Returns:
| Type | Description |
|---|---|
FeatureMaps
|
torch.Tensor: - output, attentioned pool over the last dimension (B, C, H, W) |