agrobot_base/Python/yolov8-seg/patch_model_xch.py

49 lines
1.6 KiB
Python

# 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