import torch
import torch.nn.functional as F
from core.modules.structural_gray_attention import SobelConv

def SobelLoss(attn_mask, input_img, weight=1.0):
    if input_img.shape[1] != 1:
        input_img = input_img.mean(dim=1, keepdim=True)
    sobel = SobelConv().to(input_img.device)
    edge_input = sobel(input_img)
    edge_attn = sobel(attn_mask)
    return weight * F.l1_loss(edge_attn, edge_input)