import torch import torch.nn as nn import torch.nn.functional as F class CropFCNHead(nn.Module): """ AgriFM segmentation head - faithful to GitHub implementation. Conv(embed -> embed//2) + ReLU + Conv(embed//2 -> num_classes) """ def __init__(self, embed_dim, num_classes): super().__init__() self.embed_dim = embed_dim self.num_classes = num_classes self.head = nn.Sequential( nn.Conv2d(embed_dim, embed_dim // 2, kernel_size=3, stride=1, padding=1), nn.ReLU(inplace=True), nn.Conv2d(embed_dim // 2, num_classes, kernel_size=1, stride=1, padding=0), ) def forward(self, x): """ x: (B, embed_dim, H, W) returns: (B, num_classes, H, W) """ return self.head(x)