# patch_model_xch.py import torch import torch.nn as nn def patch_yolov8_first_conv_to_xch(model_or_seq, canais: int = 5): """ Aceita: - SegmentationModel (tem .model) - nn.Sequential (já é o .model interno) """ seq = model_or_seq.model if hasattr(model_or_seq, "model") else model_or_seq # seq é tipo nn.Sequential; primeiro bloco costuma ser Conv(...) first = seq[0] # ultralytics Conv wrapper: first.conv é nn.Conv2d conv = first.conv if conv.in_channels == canais: return # já patchado old_w = conv.weight.data # [out, in, k, k] old_in = conv.in_channels out_ch, _, k1, k2 = old_w.shape new_conv = nn.Conv2d( in_channels=canais, out_channels=conv.out_channels, kernel_size=conv.kernel_size, stride=conv.stride, padding=conv.padding, dilation=conv.dilation, groups=conv.groups, bias=(conv.bias is not None), padding_mode=conv.padding_mode, ).to(conv.weight.device).type(conv.weight.dtype) with torch.no_grad(): # copia os canais existentes new_conv.weight[:, :old_in, :, :] = old_w # inicializa canais extras com a média dos 3 canais (boa heurística) if canais > old_in: mean = old_w.mean(dim=1, keepdim=True) # [out,1,k,k] new_conv.weight[:, old_in:, :, :] = mean.repeat(1, canais - old_in, 1, 1) if conv.bias is not None: new_conv.bias.copy_(conv.bias.data) # troca a conv dentro do wrapper Ultralytics first.conv = new_conv