49 lines
1.6 KiB
Python
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
|