diff --git a/.gitignore b/.gitignore index ae935a783..eccb91e00 100644 --- a/.gitignore +++ b/.gitignore @@ -93,3 +93,7 @@ AgroBase/AgroBase/bin/x64/Debug/Python/venv/ /Exemplos /Python/raspi/cam_2/imx296_pi/ /Python/raspi/cam_3/imx296_pi/ +Python/OAK/datasets/oak-fcc-3/audit/ +/Python/OAK/datasets/oak-fcc-3/benchmarks_async +/Python/OAK/datasets/oak-d/benchmark_corridor +/AgroBase/livox_visual_debugger/x64/Debug diff --git a/AgroBase/AgroBase/Forms/IHM/Operacao/Parametros/ucParametrosMapa.cs b/AgroBase/AgroBase/Forms/IHM/Operacao/Parametros/ucParametrosMapa.cs index de1e42264..1012f467e 100644 --- a/AgroBase/AgroBase/Forms/IHM/Operacao/Parametros/ucParametrosMapa.cs +++ b/AgroBase/AgroBase/Forms/IHM/Operacao/Parametros/ucParametrosMapa.cs @@ -105,7 +105,7 @@ namespace AgroBase.Forms.IHM.Operacao.Parametros _mapa.TipoMapa = op.Parametros.TipoMapa; - await _mapa.DefinirRuasSelecionadasAsync(ruasSelecionadas, false); + await _mapa.DefinirRuasSelecionadasAsync(op, ruasSelecionadas, false); lblMapaPlaceholder.Visible = false; diff --git a/AgroBase/AgroBase/Forms/IHM/Operacao/ucOperacaoNavegacao.cs b/AgroBase/AgroBase/Forms/IHM/Operacao/ucOperacaoNavegacao.cs index 3e9188490..4c36e710d 100644 --- a/AgroBase/AgroBase/Forms/IHM/Operacao/ucOperacaoNavegacao.cs +++ b/AgroBase/AgroBase/Forms/IHM/Operacao/ucOperacaoNavegacao.cs @@ -169,7 +169,7 @@ namespace AgroBase.Forms.IHM.Operacao mapa.TipoMapa = op.Parametros?.TipoMapa ?? Enums.TipoMapaOperacao.Indefinido; - await mapa.DefinirRuasSelecionadasAsync(ruasSelecionadas, false); + await mapa.DefinirRuasSelecionadasAsync(op, ruasSelecionadas, false); } else { diff --git a/AgroBase/AgroBase/Forms/IHM/Operacao/ucOperacaoParametros.cs b/AgroBase/AgroBase/Forms/IHM/Operacao/ucOperacaoParametros.cs index 8b8186250..7bbaf041a 100644 --- a/AgroBase/AgroBase/Forms/IHM/Operacao/ucOperacaoParametros.cs +++ b/AgroBase/AgroBase/Forms/IHM/Operacao/ucOperacaoParametros.cs @@ -197,7 +197,7 @@ namespace AgroBase.Forms.IHM.Operacao if (mapaAtual != null) { novaOperacao.Mapa = mapaAtual; - await mapaAtual.LimparParametrizacaoMapaAsync(); + await mapaAtual.LimparParametrizacaoMapaAsync(novaOperacao); } /* diff --git a/AgroBase/AgroBase/Forms/Operacoes/frmResultadosOperacao.Designer.cs b/AgroBase/AgroBase/Forms/Operacoes/frmResultadosOperacao.Designer.cs index f36dac0c9..602aa38c2 100644 --- a/AgroBase/AgroBase/Forms/Operacoes/frmResultadosOperacao.Designer.cs +++ b/AgroBase/AgroBase/Forms/Operacoes/frmResultadosOperacao.Designer.cs @@ -367,6 +367,11 @@ this.tabOrientacao = new System.Windows.Forms.TabPage(); this.pnlInclinacao = new System.Windows.Forms.Panel(); this.pnl3D = new System.Windows.Forms.Panel(); + this.pnlOrientacaoDados = new System.Windows.Forms.Panel(); + this.txtImuLateral = new System.Windows.Forms.TextBox(); + this.txtImuFrontal = new System.Windows.Forms.TextBox(); + this.label147 = new System.Windows.Forms.Label(); + this.label146 = new System.Windows.Forms.Label(); this.tabSonar = new System.Windows.Forms.TabPage(); this.chbSonarDeteccoes = new System.Windows.Forms.CheckBox(); this.tabControl3 = new System.Windows.Forms.TabControl(); @@ -531,11 +536,7 @@ this.label36 = new System.Windows.Forms.Label(); this.txtMomento = new System.Windows.Forms.TextBox(); this.btnPlay = new System.Windows.Forms.Button(); - this.pnlOrientacaoDados = new System.Windows.Forms.Panel(); - this.label146 = new System.Windows.Forms.Label(); - this.label147 = new System.Windows.Forms.Label(); - this.txtImuFrontal = new System.Windows.Forms.TextBox(); - this.txtImuLateral = new System.Windows.Forms.TextBox(); + this.btnCarregarMapa = new System.Windows.Forms.Button(); ((System.ComponentModel.ISupportInitialize)(this.trbMomento)).BeginInit(); this.panel1.SuspendLayout(); this.tabControl4.SuspendLayout(); @@ -590,6 +591,7 @@ ((System.ComponentModel.ISupportInitialize)(this.picRobo)).BeginInit(); this.tabOrientacao.SuspendLayout(); this.pnl3D.SuspendLayout(); + this.pnlOrientacaoDados.SuspendLayout(); this.tabSonar.SuspendLayout(); this.tabControl3.SuspendLayout(); this.tabSonarRGB.SuspendLayout(); @@ -616,7 +618,6 @@ this.panel6.SuspendLayout(); this.pnlInformacoesOperacao.SuspendLayout(); ((System.ComponentModel.ISupportInitialize)(this.picCameraCaminho)).BeginInit(); - this.pnlOrientacaoDados.SuspendLayout(); this.SuspendLayout(); // // trbMomento @@ -4131,6 +4132,51 @@ this.pnl3D.Size = new System.Drawing.Size(486, 512); this.pnl3D.TabIndex = 1; // + // pnlOrientacaoDados + // + this.pnlOrientacaoDados.Controls.Add(this.txtImuLateral); + this.pnlOrientacaoDados.Controls.Add(this.txtImuFrontal); + this.pnlOrientacaoDados.Controls.Add(this.label147); + this.pnlOrientacaoDados.Controls.Add(this.label146); + this.pnlOrientacaoDados.Location = new System.Drawing.Point(1, 362); + this.pnlOrientacaoDados.Name = "pnlOrientacaoDados"; + this.pnlOrientacaoDados.Size = new System.Drawing.Size(150, 150); + this.pnlOrientacaoDados.TabIndex = 0; + // + // txtImuLateral + // + this.txtImuLateral.Location = new System.Drawing.Point(6, 87); + this.txtImuLateral.Name = "txtImuLateral"; + this.txtImuLateral.Size = new System.Drawing.Size(100, 20); + this.txtImuLateral.TabIndex = 3; + this.txtImuLateral.TextAlign = System.Windows.Forms.HorizontalAlignment.Right; + // + // txtImuFrontal + // + this.txtImuFrontal.Location = new System.Drawing.Point(6, 32); + this.txtImuFrontal.Name = "txtImuFrontal"; + this.txtImuFrontal.Size = new System.Drawing.Size(100, 20); + this.txtImuFrontal.TabIndex = 1; + this.txtImuFrontal.TextAlign = System.Windows.Forms.HorizontalAlignment.Right; + // + // label147 + // + this.label147.AutoSize = true; + this.label147.Location = new System.Drawing.Point(3, 71); + this.label147.Name = "label147"; + this.label147.Size = new System.Drawing.Size(39, 13); + this.label147.TabIndex = 2; + this.label147.Text = "Lateral"; + // + // label146 + // + this.label146.AutoSize = true; + this.label146.Location = new System.Drawing.Point(3, 16); + this.label146.Name = "label146"; + this.label146.Size = new System.Drawing.Size(39, 13); + this.label146.TabIndex = 1; + this.label146.Text = "Frontal"; + // // tabSonar // this.tabSonar.Controls.Add(this.chbSonarDeteccoes); @@ -5875,10 +5921,10 @@ // // btnCarregarOperacao // - this.btnCarregarOperacao.Location = new System.Drawing.Point(1180, 587); + this.btnCarregarOperacao.Location = new System.Drawing.Point(1138, 587); this.btnCarregarOperacao.Margin = new System.Windows.Forms.Padding(2); this.btnCarregarOperacao.Name = "btnCarregarOperacao"; - this.btnCarregarOperacao.Size = new System.Drawing.Size(128, 28); + this.btnCarregarOperacao.Size = new System.Drawing.Size(115, 28); this.btnCarregarOperacao.TabIndex = 4; this.btnCarregarOperacao.Text = "Carregar Opereação"; this.btnCarregarOperacao.UseVisualStyleBackColor = true; @@ -5907,7 +5953,7 @@ // // btnPlay // - this.btnPlay.Location = new System.Drawing.Point(1120, 587); + this.btnPlay.Location = new System.Drawing.Point(1078, 587); this.btnPlay.Margin = new System.Windows.Forms.Padding(2); this.btnPlay.Name = "btnPlay"; this.btnPlay.Size = new System.Drawing.Size(56, 28); @@ -5916,56 +5962,23 @@ this.btnPlay.UseVisualStyleBackColor = true; this.btnPlay.Click += new System.EventHandler(this.btnPlay_Click); // - // pnlOrientacaoDados + // btnCarregarMapa // - this.pnlOrientacaoDados.Controls.Add(this.txtImuLateral); - this.pnlOrientacaoDados.Controls.Add(this.txtImuFrontal); - this.pnlOrientacaoDados.Controls.Add(this.label147); - this.pnlOrientacaoDados.Controls.Add(this.label146); - this.pnlOrientacaoDados.Location = new System.Drawing.Point(1, 362); - this.pnlOrientacaoDados.Name = "pnlOrientacaoDados"; - this.pnlOrientacaoDados.Size = new System.Drawing.Size(150, 150); - this.pnlOrientacaoDados.TabIndex = 0; - // - // label146 - // - this.label146.AutoSize = true; - this.label146.Location = new System.Drawing.Point(3, 16); - this.label146.Name = "label146"; - this.label146.Size = new System.Drawing.Size(39, 13); - this.label146.TabIndex = 1; - this.label146.Text = "Frontal"; - // - // label147 - // - this.label147.AutoSize = true; - this.label147.Location = new System.Drawing.Point(3, 71); - this.label147.Name = "label147"; - this.label147.Size = new System.Drawing.Size(39, 13); - this.label147.TabIndex = 2; - this.label147.Text = "Lateral"; - // - // txtImuFrontal - // - this.txtImuFrontal.Location = new System.Drawing.Point(6, 32); - this.txtImuFrontal.Name = "txtImuFrontal"; - this.txtImuFrontal.Size = new System.Drawing.Size(100, 20); - this.txtImuFrontal.TabIndex = 1; - this.txtImuFrontal.TextAlign = System.Windows.Forms.HorizontalAlignment.Right; - // - // txtImuLateral - // - this.txtImuLateral.Location = new System.Drawing.Point(6, 87); - this.txtImuLateral.Name = "txtImuLateral"; - this.txtImuLateral.Size = new System.Drawing.Size(100, 20); - this.txtImuLateral.TabIndex = 3; - this.txtImuLateral.TextAlign = System.Windows.Forms.HorizontalAlignment.Right; + this.btnCarregarMapa.Location = new System.Drawing.Point(1257, 587); + this.btnCarregarMapa.Margin = new System.Windows.Forms.Padding(2); + this.btnCarregarMapa.Name = "btnCarregarMapa"; + this.btnCarregarMapa.Size = new System.Drawing.Size(77, 28); + this.btnCarregarMapa.TabIndex = 96; + this.btnCarregarMapa.Text = "Mapa"; + this.btnCarregarMapa.UseVisualStyleBackColor = true; + this.btnCarregarMapa.Click += new System.EventHandler(this.btnCarregarMapa_Click); // // frmResultadosOperacao // this.AutoScaleDimensions = new System.Drawing.SizeF(6F, 13F); this.AutoScaleMode = System.Windows.Forms.AutoScaleMode.Font; this.ClientSize = new System.Drawing.Size(1345, 632); + this.Controls.Add(this.btnCarregarMapa); this.Controls.Add(this.btnPlay); this.Controls.Add(this.txtMomento); this.Controls.Add(this.label36); @@ -6061,6 +6074,8 @@ ((System.ComponentModel.ISupportInitialize)(this.picRobo)).EndInit(); this.tabOrientacao.ResumeLayout(false); this.pnl3D.ResumeLayout(false); + this.pnlOrientacaoDados.ResumeLayout(false); + this.pnlOrientacaoDados.PerformLayout(); this.tabSonar.ResumeLayout(false); this.tabSonar.PerformLayout(); this.tabControl3.ResumeLayout(false); @@ -6096,8 +6111,6 @@ this.pnlInformacoesOperacao.ResumeLayout(false); this.pnlInformacoesOperacao.PerformLayout(); ((System.ComponentModel.ISupportInitialize)(this.picCameraCaminho)).EndInit(); - this.pnlOrientacaoDados.ResumeLayout(false); - this.pnlOrientacaoDados.PerformLayout(); this.ResumeLayout(false); this.PerformLayout(); @@ -6588,5 +6601,6 @@ private System.Windows.Forms.TextBox txtImuFrontal; private System.Windows.Forms.Label label147; private System.Windows.Forms.Label label146; + private System.Windows.Forms.Button btnCarregarMapa; } } \ No newline at end of file diff --git a/AgroBase/AgroBase/Forms/Operacoes/frmResultadosOperacao.cs b/AgroBase/AgroBase/Forms/Operacoes/frmResultadosOperacao.cs index 9be4a4842..bd4445d75 100644 --- a/AgroBase/AgroBase/Forms/Operacoes/frmResultadosOperacao.cs +++ b/AgroBase/AgroBase/Forms/Operacoes/frmResultadosOperacao.cs @@ -1,5 +1,6 @@ using AgroBase.Models; using AgroBase.Models.Modules; +using AgroBase.Models.Operacoes; using AgroBase.Models.Operadores; using AgroBase.Properties; using AgroBase.Services; @@ -12,6 +13,7 @@ using System.Drawing.Drawing2D; using System.Globalization; using System.IO; using System.Linq; +using System.Text; using System.Threading.Tasks; using System.Windows.Forms; @@ -53,9 +55,14 @@ namespace AgroBase.Forms.Operacoes DateTime MomentoAtual = DateTime.Now; AsyncTaskTimerModel tmrPlay; + + private List> _ruasMapaReplay = new List>(); + private List _trajetoriaFixaReplay = new List(); + MapasModel Mapa = new MapasModel(); - private Visualizador3DService visualizador3D; + //TrajetoriaMapaOperacaoModel Trajetoria = new TrajetoriaMapaOperacaoModel(new List>()); private MapaDinamicoModel MapaDinamico; + private Visualizador3DService visualizador3D; private int idxMomentoAtual { get @@ -165,12 +172,46 @@ namespace AgroBase.Forms.Operacoes LogsOperacao = JsonConvert.DeserializeObject(OperacaoModel.DeserializarDadosOperacao(ofd.FileName, data.operacao)).ToList(); LogsGPS = JsonConvert.DeserializeObject(OperacaoModel.DeserializarDadosOperacao(ofd.FileName, Enums.T_Code.Gps.ToString())).ToList(); - GPSModel pos = GPSService.UltimaLeitura.Clone(); - GPSService.UltimaLeitura = LogsGPS.FirstOrDefault(); - bool sucesso = op.CarregarParametrizacaoOperacao(null, null, ofd.FileName.Replace("data.lgop", "operacao.opr")); - bool condicaoOperacaoCarregada() => op.Trajetoria?.CorredorAtual != null; - bool resultado = await FuncoesGlobais.AguardarCondicaoAsync(condicaoOperacaoCarregada, 15000, 200); - GPSService.UltimaLeitura = pos; + + string caminhoOpr = Path.Combine( + Path.GetDirectoryName(ofd.FileName), + "operacao.opr" + ); + + OperacaoParametrosModel parametrosOperacao = null; + + if (File.Exists(caminhoOpr)) + { + parametrosOperacao = + JsonConvert.DeserializeObject( + File.ReadAllText(caminhoOpr, Encoding.UTF8) + ); + + if (parametrosOperacao != null) + { + GPSModel posAnterior = + GPSService.UltimaLeitura?.Clone(); + + try + { + // Importantíssimo: + // a trajetória será projetada tomando como referência + // a posição inicial real daquela operação. + GPSService.UltimaLeitura = + LogsGPS.FirstOrDefault()?.Clone(); + + Variaveis.OperacaoEmAndamento = OperacaoModel.CarregarParamerosOperacaoBase( + parametrosOperacao + ); + } + finally + { + if (posAnterior != null) + GPSService.UltimaLeitura = posAnterior; + } + } + } + LogsControle = JsonConvert.DeserializeObject(OperacaoModel.DeserializarDadosOperacao(ofd.FileName, "controle")).ToList(); LogsCooler = JsonConvert.DeserializeObject(OperacaoModel.DeserializarDadosOperacao(ofd.FileName, "cooler")).ToList(); @@ -239,7 +280,7 @@ namespace AgroBase.Forms.Operacoes }; MapasService.CriarArquivoMapa(new List>() { LogsGPS.ToList() }, 1517, data.id); Mapa.btnCarregar_Click(sender, e, Path.Combine(Variaveis.CaminhoSistema, Variaveis.CaminhoMapasConvertidos + data.id + ".json")); - + trbMomento.Maximum = LogsOperacao.Count - 1; trbMomento.Minimum = 0; trbMomento.Value = 0; @@ -1309,28 +1350,60 @@ namespace AgroBase.Forms.Operacoes pnlOrientacaoGPS.Invalidate(); - var _trajaetoria = LogsGPS.Select(x => new PontoTrajetoriaModel(Enums.TipoPontoRua.Indefinido) - { - Posicao = x - }) - .ToList(); - - var LogOpe = LogsOperacao[idxMomentoAtual]; var LogSnr = LogsVisualWorker[idxMomentoAtual]; var LogTrj = LogsTrajetoria[idxMomentoAtual]; var LogCtrl = LogsControle[idxMomentoAtual]; + + var trajetoriaFixa = + (_trajetoriaFixaReplay?.Count ?? 0) > 0 + ? _trajetoriaFixaReplay + : Variaveis.OperacaoEmAndamento + ?.Trajetoria + ?._TrajetoriaFixa + ?? new List(); + + var trajetoriaDinamica = + MontarTrajetoriaDinamicaReplay( + Log, + LogTrj + ); + + var ruasMapa = + (_ruasMapaReplay?.Count ?? 0) > 0 + ? _ruasMapaReplay + : Variaveis.OperacaoEmAndamento + ?.Mapa + ?.TrajetoriaMapa + ?? new List>(); + + var rastroRover = + LogsGPS + .Take(idxMomentoAtual + 1) + .ToList(); + + + Variaveis.OperacaoEmAndamento.Trajetoria.LoopAtualizaDados(); MapaDinamico.AtualizarDados( - //(LogSnr?.Iniciado ?? false) && LogSnr.Resumo.DirecaoDesvio != Enums.Direcao.Frente ? LogSnr.Resumo.ObstaculoCritico : null, - null, + null, (float)LogTrj.AnguloCaminho, (float)Log.AnguloCarroDefinido, Log, - _trajaetoria, - Variaveis.OperacaoEmAndamento.Trajetoria?._TrajetoriaFixa, - new List(LogsGPS.Take(LogsGPS.IndexOf(Log) + 1)), - Variaveis.OperacaoEmAndamento.Mapa.TrajetoriaMapa, + + // Roxo: trajetória ainda restante naquele instante + trajetoriaDinamica, + + // Verde/azul: trajetória fixa completa projetada + trajetoriaFixa, + + // Laranja: caminho realmente percorrido até aquele instante + rastroRover, + + // Azul/preto: linhas originais do mapa + ruasMapa, + + // Verde tracejado: simulação MPC daquele frame LogCtrl.SimulacaoMPC ); } @@ -2334,6 +2407,48 @@ namespace AgroBase.Forms.Operacoes } } } + + private async void btnCarregarMapa_Click(object sender, EventArgs e) + { + + } + + private List MontarTrajetoriaDinamicaReplay(GPSModel posicaoAtual, OperacaoSensoriamentoLogTrajetoriaModel logTrajetoria) + { + var resultado = new List(); + + if (posicaoAtual != null) + { + resultado.Add( + new PontoTrajetoriaModel(Enums.TipoPontoRua.PosicaoRobo) + { + Posicao = posicaoAtual, + LarguraCorredor = 1.0, + idxCorredor = logTrajetoria?.CorredorAtual?.Idx ?? 0, + Visitado = true + } + ); + } + + if (_trajetoriaFixaReplay == null || + _trajetoriaFixaReplay.Count == 0) + { + return resultado; + } + + int idxProximo = + logTrajetoria?.ProximoPonto?.idxPonto ?? 0; + + resultado.AddRange( + _trajetoriaFixaReplay + .Where(x => + x != null && + x.idxPonto >= idxProximo) + ); + + return resultado; + } + } } diff --git a/AgroBase/AgroBase/Forms/Simuladores/frmTreinamentoIA.cs b/AgroBase/AgroBase/Forms/Simuladores/frmTreinamentoIA.cs index d34104928..9d220bddc 100644 --- a/AgroBase/AgroBase/Forms/Simuladores/frmTreinamentoIA.cs +++ b/AgroBase/AgroBase/Forms/Simuladores/frmTreinamentoIA.cs @@ -88,7 +88,7 @@ namespace AgroBase.Forms txtAnguloGPS.Text = anguloInicial.ToString("0.00"); txtDistanciaGPS.Text = distanciaInicial.ToString("0.00"); Variaveis.OperacaoEmAndamento.GPSTrajetoria = new System.Collections.Generic.List(); - Variaveis.OperacaoEmAndamento.Trajetoria = new TrajetoriaMapaOperacaoModel(Variaveis.OperacaoEmAndamento.Mapa.TrajetoriaMapa); + Variaveis.OperacaoEmAndamento.Trajetoria = new TrajetoriaMapaOperacaoModel(Variaveis.OperacaoEmAndamento, Variaveis.OperacaoEmAndamento.Mapa.TrajetoriaMapa); Variaveis.OperacaoEmAndamento.Trajetoria.ProjetarTrajetoriaFixa(); _robotState = new TreinamentoIAModel(); } diff --git a/AgroBase/AgroBase/Models/MapasModel.cs b/AgroBase/AgroBase/Models/MapasModel.cs index ae9cb907a..8b83102db 100644 --- a/AgroBase/AgroBase/Models/MapasModel.cs +++ b/AgroBase/AgroBase/Models/MapasModel.cs @@ -344,7 +344,7 @@ namespace AgroBase.Models _mapaEnviado = true; } - public void DefinirRuasSelecionadas(IEnumerable ruas, bool popularTrajetoria = false) + public void DefinirRuasSelecionadas(OperacaoModel op, IEnumerable ruas, bool popularTrajetoria = false) { RuasPercorrer = ruas == null @@ -361,12 +361,12 @@ namespace AgroBase.Models } if (popularTrajetoria) - PopularTrajetoriaMapa(); + PopularTrajetoriaMapa(op); } - public async Task DefinirRuasSelecionadasAsync(IEnumerable ruas, bool popularTrajetoria = false) + public async Task DefinirRuasSelecionadasAsync(OperacaoModel op, IEnumerable ruas, bool popularTrajetoria = false) { - DefinirRuasSelecionadas(ruas, popularTrajetoria); + DefinirRuasSelecionadas(op, ruas, popularTrajetoria); if (!_paginaCarregada || browser?.CoreWebView2 == null) { @@ -378,11 +378,11 @@ namespace AgroBase.Models await browser.CoreWebView2.ExecuteScriptAsync("definirSelecaoRuas(" + ruasJson + ");"); } - public Task DefinirRuasSelecionadasAsync(IEnumerable ruas) + public Task DefinirRuasSelecionadasAsync(OperacaoModel op, IEnumerable ruas) { IEnumerable convertidas = ruas == null ? null : ruas.Select(x => x.ToString()); - return DefinirRuasSelecionadasAsync(convertidas); + return DefinirRuasSelecionadasAsync(op, convertidas); } private bool PodeSelecionarRuas() @@ -424,7 +424,7 @@ namespace AgroBase.Models .Distinct() .ToList(); - PopularTrajetoriaMapa(); + PopularTrajetoriaMapa(Variaveis.OperacaoEmAndamento); RuasSelecionadasAlteradas?.Invoke(this, EventArgs.Empty); } catch (Exception ex) @@ -433,7 +433,7 @@ namespace AgroBase.Models } } - public void PopularTrajetoriaMapa(MapaFeatureCollectionModel dados = null) + public void PopularTrajetoriaMapa(OperacaoModel op, MapaFeatureCollectionModel dados = null) { if (dados != null) mapaService.DefinirMapa(dados, mapaService.NomeArquivos ?? "Mapa"); @@ -512,14 +512,12 @@ namespace AgroBase.Models ruasMontadas.Add(rua); } - OperacaoModel op = Variaveis.OperacaoEmAndamento; - if (op == null) throw new InvalidOperationException("Operação indisponível para instalar a trajetória."); var trajetoriaAnterior = op.Trajetoria; var mapaAnterior = TrajetoriaMapa; - var novaTrajetoria = new TrajetoriaMapaOperacaoModel(ruasMontadas); + var novaTrajetoria = new TrajetoriaMapaOperacaoModel(op, ruasMontadas); try { @@ -647,9 +645,9 @@ namespace AgroBase.Models pnlMapa = null; } - public async Task LimparParametrizacaoMapaAsync() + public async Task LimparParametrizacaoMapaAsync(OperacaoModel op) { - DefinirRuasSelecionadas(new List(), false); + DefinirRuasSelecionadas(op, new List(), false); if (_paginaCarregada && browser?.CoreWebView2 != null) { await browser.CoreWebView2.ExecuteScriptAsync("definirSelecaoRuas([]);"); diff --git a/AgroBase/AgroBase/Models/Operacoes/OperacaoModel.cs b/AgroBase/AgroBase/Models/Operacoes/OperacaoModel.cs index 8686b9426..9e6c6ebee 100644 --- a/AgroBase/AgroBase/Models/Operacoes/OperacaoModel.cs +++ b/AgroBase/AgroBase/Models/Operacoes/OperacaoModel.cs @@ -476,7 +476,6 @@ namespace AgroBase.Models Simulando = false, GPSTrajetoria = new List(), Mapa = new MapasModel(), - Trajetoria = new TrajetoriaMapaOperacaoModel(new List>()), ControleAnterior = new OperacaoControleModel(), Parametros = new OperacaoParametrosModel() { @@ -499,10 +498,12 @@ namespace AgroBase.Models } } }; + op.Trajetoria = new TrajetoriaMapaOperacaoModel(op, new List>()); op.Sensoriamento = new OperacaoSensoriamentoConjuntoModel(op) { Operacao = new OperacaoSensoriamentoLogModel() { + Modo = Modo, OperacaoIniciada = false, DataInicio = DateTime.MinValue, DataFim = DateTime.MinValue @@ -559,7 +560,7 @@ namespace AgroBase.Models AtuPercentualErvasBicoOff = 2, AtuPercentualErvasBicoOn = 1, AtuPercentualInicioPulverizacao = 85, - AtuPressaoLinha = 18, + AtuPressaoLinha = 22, AtuAgitadorModo = ModoAgitadorCalda.SemAgitacao, AtuModoControle = ModoControleBomba.PID, AtuModeloCabeca = CabecasModelo.Target, @@ -666,14 +667,14 @@ namespace AgroBase.Models DirVelocidadeMovimento = 80, DirTipoMovimento = TiposControladorDirecional.MPC, - DistanciaManobra = 3.0, + DistanciaManobra = 2.2, AtuAlturaAreaPulverizacao = 10, AtuDuracaoAtuacao = 0, AtuPercentualErvasBicoOff = 2, AtuPercentualErvasBicoOn = 1, AtuPercentualInicioPulverizacao = 85, - AtuPressaoLinha = 18, + AtuPressaoLinha = 22, AtuAgitadorModo = ModoAgitadorCalda.Continuo, AtuModoControle = ModoControleBomba.PID, AtuModeloCabeca = CabecasModelo.Target, @@ -1312,9 +1313,9 @@ namespace AgroBase.Models if (parametros.Mapa != null) { novaOp.Mapa = new MapasModel(); - novaOp.Mapa.DefinirRuasSelecionadas(ruasSelecionadas); + novaOp.Mapa.DefinirRuasSelecionadas(novaOp, ruasSelecionadas); novaOp.Mapa.TipoMapa = novaOp.Parametros.TipoMapa; - novaOp.Mapa.PopularTrajetoriaMapa(parametros.Mapa); + novaOp.Mapa.PopularTrajetoriaMapa(novaOp, parametros.Mapa); } novaOp?.DispAtu?.Dados?.PrepararAtuadorParaOperacao(); diff --git a/AgroBase/AgroBase/Models/TrajetoriaMapaOperacaoModel.cs b/AgroBase/AgroBase/Models/TrajetoriaMapaOperacaoModel.cs index 1b9d1c39e..a539806c2 100644 --- a/AgroBase/AgroBase/Models/TrajetoriaMapaOperacaoModel.cs +++ b/AgroBase/AgroBase/Models/TrajetoriaMapaOperacaoModel.cs @@ -16,8 +16,9 @@ namespace AgroBase.Models { public class TrajetoriaMapaOperacaoModel { - public TrajetoriaMapaOperacaoModel(List> RuasMapa) + public TrajetoriaMapaOperacaoModel(OperacaoModel _op, List> RuasMapa) { + op = _op; RuasPlantacao = ClonarRuasMapa(RuasMapa); AutonomiaCorredor = new AutonomiaCorredorModel(); } @@ -47,12 +48,15 @@ namespace AgroBase.Models private const double ToleranciaCoordenada = 1e-12; private const double DistanciaMaximaSaltoMapaM = 25.0; + [JsonIgnore] + private readonly OperacaoModel op; + #region PARAMETROS public double AnguloAberturaCurva { get; set; } = 25; // Angulo usado para deslocar o ponto de curva public double DistanciaProjecaoRua => (VariaveisEquipamento.DistanciaEntreEixosCm / 100.0 / 2.0) + - (Variaveis.OperacaoEmAndamento?.Parametros?.Controle?.DistanciaManobra ?? 3.0); // Distancia para projetar o primeiro ponto para fora do corredor + (op?.Parametros?.Controle?.DistanciaManobra ?? 3.0); // Distancia para projetar o primeiro ponto para fora do corredor public static double DistanciaEntrePontos { get; set; } = 0.8; // Distancia entre os pontos dentro do corredor public static double DistanciaEntrePontosCurva { get; set; } = 0.25; // Distancia entre os pontos durante a curva entre corredores private double DistanciaManobraEntreRuas { get; set; } = 3.0; // Distancia máxima para gerar a curva de conexão entre os corredores @@ -318,7 +322,7 @@ namespace AgroBase.Models public static double DistanciaMaximaEntreLeituras { get; private set; } public void AtualizarDistanciaMaximaEntreLeituras() { - double percentualVelocidade = Variaveis.OperacaoEmAndamento?.Controle?.PercentualVelocidadeSP ?? 0; + double percentualVelocidade = op?.Controle?.PercentualVelocidadeSP ?? 0; double velocidadeCarroMs = FuncoesMatematicas.CalculaVelocidadeMsPercentual(percentualVelocidade); double distanciaMaxima = velocidadeCarroMs / Math.Max(1, GPSService.TaxaAmostragemHz); @@ -673,7 +677,7 @@ namespace AgroBase.Models if (CorredorAtual != null) { bool operacaoEmAndamento = - Variaveis.OperacaoEmAndamento?.Sensoriamento?.Operacao?.StatusOperacaoAtual == + op?.Sensoriamento?.Operacao?.StatusOperacaoAtual == StatusOperacao.EmAndamento; if ((RetornandoBase && CorredorAtual.Dentro) || (!RetornandoBase && operacaoEmAndamento)) @@ -784,8 +788,6 @@ namespace AgroBase.Models public string TempoEstimadoRestante { get; private set; } public void AtualizarTempoEstimado(double? velocidadeSemErvasMs = null, double? velocidadeComErvasMs = null) { - var op = Variaveis.OperacaoEmAndamento; - if (op?.Sensoriamento == null || op?.Parametros?.Controle == null) { TempoEstimadoOperacao = "00:00:00"; @@ -871,8 +873,6 @@ namespace AgroBase.Models if (CorredorAtual.Idx > 0) return; if (!CorredorAtual.Dentro) return; - var op = Variaveis.OperacaoEmAndamento; - if (op?.Sensoriamento == null) return; @@ -952,7 +952,7 @@ namespace AgroBase.Models { MarcarVisitadoAte(_TrajetoriaFixa.Count - 1); - Variaveis.OperacaoEmAndamento.Sensoriamento?.InserirLog( + op.Sensoriamento?.InserirLog( T_Code.Trj, StatusModulo.Operante, 100, @@ -976,7 +976,6 @@ namespace AgroBase.Models } private void AtualizarErroCombinado() { - var op = Variaveis.OperacaoEmAndamento; var gps = GPSPosicaoAtual; if (op?.Parametros?.Controle == null || gps == null || PontoAtual?.Posicao == null || ProximoPonto?.Posicao == null) @@ -1139,6 +1138,10 @@ namespace AgroBase.Models for (int i = idxInicial; i <= idxFinal; i++) { + if (PontoAtual.idxPonto == _TrajetoriaFixa[i].idxPonto) + { + + } _TrajetoriaFixa[i].AtualizarPropriedades( _gpsAtualCiclo, _gpsAnteriorCiclo, @@ -1261,7 +1264,7 @@ namespace AgroBase.Models if (!confirmou) return; - Variaveis.OperacaoEmAndamento.Sensoriamento?.InserirLog( + op.Sensoriamento?.InserirLog( T_Code.Trj, StatusModulo.Operante, 100, @@ -1468,7 +1471,7 @@ namespace AgroBase.Models private bool VerificaInicioOperacaoMeioRua(bool considerarStatus) { - if (_TrajetoriaFixaDefinida && !VerificacaoInicialMeioRuaConcluida && ((considerarStatus && new List() { StatusOperacao.Aguardando, StatusOperacao.EmAndamento }.Contains(Variaveis.OperacaoEmAndamento.Sensoriamento.Operacao.StatusOperacaoAtual)) || !considerarStatus)) + if (_TrajetoriaFixaDefinida && !VerificacaoInicialMeioRuaConcluida && ((considerarStatus && new List() { StatusOperacao.Aguardando, StatusOperacao.EmAndamento }.Contains(op.Sensoriamento.Operacao.StatusOperacaoAtual)) || !considerarStatus)) { var _posicaoAtual = GPSPosicaoAtual; if (_posicaoAtual == null) @@ -1487,7 +1490,7 @@ namespace AgroBase.Models if (!double.IsNaN(distanciaUltimoPonto) && !double.IsInfinity(distanciaUltimoPonto) && distanciaUltimoPonto < 5.0) { msg = $"Início no meio da rua ignorado: robô está a {distanciaUltimoPonto:0.00}m do fim. Nova operação começará do zero."; - Variaveis.OperacaoEmAndamento?.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Operante, 100, msg); + op?.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Operante, 100, msg); Variaveis.MostrarLog($"[TRJ] {msg}"); return false; @@ -1504,7 +1507,7 @@ namespace AgroBase.Models ) { msg = "Inicio no meio da rua rejeitado: geometria e trajetoria discordam sobre o corredor atual."; - Variaveis.OperacaoEmAndamento?.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Falha, 0, msg); + op?.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Falha, 0, msg); Variaveis.MostrarLog("[TRJ] " + msg); throw new InvalidOperationException(msg); } @@ -1517,14 +1520,14 @@ namespace AgroBase.Models CorredorAtual?.AtualizarDados(); msg = $"Início no meio da rua confirmado: idx={idxPontoMaisProximo}/{_TrajetoriaFixa.Count - 1}, marcados={idxPontoMaisProximo + 1}/{_TrajetoriaFixa.Count}."; - Variaveis.OperacaoEmAndamento?.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Alerta, 100, msg); + op?.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Alerta, 100, msg); Variaveis.MostrarLog($"[TRJ] {msg}"); return true; } msg = $"Início no meio da rua ignorado: robô está fora do corredor. Nova operação começará do zero."; - Variaveis.OperacaoEmAndamento?.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Operante, 100, msg); + op?.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Operante, 100, msg); Variaveis.MostrarLog($"[TRJ] {msg}"); } @@ -1534,7 +1537,6 @@ namespace AgroBase.Models private void VerificaFimCorredorSegmentacao() { - var op = Variaveis.OperacaoEmAndamento; double limiteFim = op?.Parametros?.Controle?.AnteciparManobraCorredorM ?? 0; if ( @@ -1972,7 +1974,7 @@ namespace AgroBase.Models catch (Exception ex) { Variaveis.MostrarLog("[TRJ] Falha no ciclo da trajetoria: " + ex.Message); - Variaveis.OperacaoEmAndamento?.Sensoriamento?.InserirLog( + op?.Sensoriamento?.InserirLog( T_Code.Trj, StatusModulo.Falha, 0, @@ -2005,8 +2007,6 @@ namespace AgroBase.Models private void LoopAtualizaDadosCore() { - var op = Variaveis.OperacaoEmAndamento; - if (!(op?.Parametros?.ControleAutomatico ?? false)) return; @@ -2023,8 +2023,6 @@ namespace AgroBase.Models private void AtualizarDadosControleRedis() { - var op = Variaveis.OperacaoEmAndamento; - if (op == null || !op.Sensoriamento.Operacao.OperacaoIniciada) { return; @@ -2270,7 +2268,7 @@ namespace AgroBase.Models DefinirPontoAtual(); DefinirProximoPonto(); - Variaveis.OperacaoEmAndamento.Sensoriamento?.InserirLog( + op.Sensoriamento?.InserirLog( T_Code.Trj, StatusModulo.Operante, 100, @@ -3203,7 +3201,7 @@ namespace AgroBase.Models List> ruasParaProjetar; - if (Variaveis.OperacaoEmAndamento?.Mapa?.TipoMapa == TipoMapaOperacao.Corredores) + if (op?.Mapa?.TipoMapa == TipoMapaOperacao.Corredores) { throw new NotSupportedException( "O tipo de mapa Corredores ainda nao possui geracao validada para campo. " + @@ -3679,7 +3677,7 @@ namespace AgroBase.Models private void ProjetarTrajetoriaFixaCore(GPSModel posicaoRobo) { - if (Variaveis.OperacaoEmAndamento?.Sensoriamento?.Operacao?.Modo != ModoOperacao.MapaGPS) + if (op?.Parametros?.Modo != ModoOperacao.MapaGPS) return; List _trajetoriaFixa = new List(); @@ -3693,7 +3691,7 @@ namespace AgroBase.Models catch (Exception ex) { Variaveis.MostrarLog($"[TRJ] Trajetória rejeitada: {ex.Message}"); - Variaveis.OperacaoEmAndamento?.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Falha, 0, $"Trajetória rejeitada: {ex.Message}"); + op?.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Falha, 0, $"Trajetória rejeitada: {ex.Message}"); throw; } @@ -3810,7 +3808,8 @@ namespace AgroBase.Models Corredores[idx] = CorredorAtual; double anguloProjetar1 = GPSUtils.CalcularOrientacao(CorredorAtual.Skip(1).FirstOrDefault(), CorredorAtual.FirstOrDefault()); - GPSModel PrimeiroPonto = GPSUtils.ProjetarPontoDeslocado(CorredorAtual.FirstOrDefault(), DistanciaProjecaoRua * 0.8, anguloProjetar1); + double percentualAcrescimoPontoAberturaCurvaPrimeiroPonto = 0.7; + GPSModel PrimeiroPonto = GPSUtils.ProjetarPontoDeslocado(CorredorAtual.FirstOrDefault(), DistanciaProjecaoRua * percentualAcrescimoPontoAberturaCurvaPrimeiroPonto, anguloProjetar1); var ultimoPontoTrajetoria = _trajetoriaFixa.LastOrDefault(); @@ -3823,7 +3822,7 @@ namespace AgroBase.Models Direcao = direcaoAtual, Orientacao = GPSUtils.CalcularOrientacao(ultimoPontoTrajetoria.Posicao, PrimeiroPonto), Visitado = false, - LarguraCorredor = larguraCorredorMenor + LarguraCorredor = LarguraCorredorPadrao // larguraCorredorMenor }; _trajetoriaFixa.Add(PontoInicial); @@ -3849,8 +3848,8 @@ namespace AgroBase.Models _trajetoriaFixa.Add(PontoTrajetoria); } - double percentualAcrescimoPontoAberturaCurva = 0.8; - double distancia_projetar = !ultimoCorredor ? (DistanciaProjecaoRua * percentualAcrescimoPontoAberturaCurva) : DistanciaProjecaoRua; + double percentualAcrescimoPontoAberturaCurvaUltimoPonto = 0.8; + double distancia_projetar = !ultimoCorredor ? (DistanciaProjecaoRua * percentualAcrescimoPontoAberturaCurvaUltimoPonto) : DistanciaProjecaoRua; GPSModel ultimoPontoCorredorAtual = CorredorAtual.Last(); double anguloProjetar2 = GPSUtils.CalcularOrientacao(CorredorAtual.Skip(CorredorAtual.Count() - 2).FirstOrDefault(), ultimoPontoCorredorAtual); @@ -3882,7 +3881,7 @@ namespace AgroBase.Models Direcao = direcaoAtual, Orientacao = GPSUtils.CalcularOrientacao(_trajetoriaFixa.Last().Posicao, UltimoPonto), Visitado = false, - LarguraCorredor = larguraCorredorMenor // ultimoCorredor ? larguraCorredorMenor : (larguraCorredorMenor * 0.8) + LarguraCorredor = LarguraCorredorPadrao * 1.1 // larguraCorredorMenor // ultimoCorredor ? larguraCorredorMenor : (larguraCorredorMenor * 0.8) }; _trajetoriaFixa.Add(PontoFinal); @@ -4376,7 +4375,6 @@ namespace AgroBase.Models double distanciaTotalAnterior = DistanciaTotal; double distanciaPercorridaAnterior = DistanciaPercorrida; var autonomiaAnterior = AutonomiaCorredor?.Clone(); - var op = Variaveis.OperacaoEmAndamento; var parametrosAnteriores = op?.Parametros; bool simulandoAnterior = op?.Simulando ?? false; bool? operacaoIniciadaAnterior = op?.Sensoriamento?.Operacao?.OperacaoIniciada; @@ -4407,7 +4405,7 @@ namespace AgroBase.Models } Variaveis.MostrarLog("[RETORNO] Trajetoria rejeitada: " + ex.Message); - Variaveis.OperacaoEmAndamento?.Sensoriamento?.InserirLog( + op?.Sensoriamento?.InserirLog( T_Code.Trj, StatusModulo.Falha, 0, @@ -4649,7 +4647,7 @@ namespace AgroBase.Models AtualizarTrajetoriaDinamicaCore(); AtualizarDistanciaRestante(); - Variaveis.OperacaoEmAndamento.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Operante, 100, $"[RETORNO] Trajetória aplicada ({motivo}) - Pontos: {traj.Count}"); + op.Sensoriamento?.InserirLog(T_Code.Trj, StatusModulo.Operante, 100, $"[RETORNO] Trajetória aplicada ({motivo}) - Pontos: {traj.Count}"); } private List CortarFinalDaRota(List pontos, double distanciaAntesDoFimM) @@ -4689,8 +4687,6 @@ namespace AgroBase.Models private void AtualizarDadosOperacaoParaRetorno(List waypointsManual) { - var op = Variaveis.OperacaoEmAndamento; - if (op == null) throw new InvalidOperationException("Operacao indisponivel ao iniciar retorno."); @@ -4784,7 +4780,7 @@ namespace AgroBase.Models */ if (!ReferenceEquals( op, - Variaveis.OperacaoEmAndamento)) + op)) { Variaveis.MostrarLog( "[RETORNO/MAPSYNC] Operação mudou durante sincronização." @@ -5971,7 +5967,7 @@ namespace AgroBase.Models */ limiteDistancia = Math.Min(limiteDistancia, 0.80); - if (NaMargem && DistanciaAtual <= limiteDistancia) + if (NaMargem) // && DistanciaAtual <= limiteDistancia { Visitado = true; } diff --git a/AgroBase/AgroBase/Services/Operadores/HealthWorkerService.cs b/AgroBase/AgroBase/Services/Operadores/HealthWorkerService.cs index 14b8bac52..b5da57be5 100644 --- a/AgroBase/AgroBase/Services/Operadores/HealthWorkerService.cs +++ b/AgroBase/AgroBase/Services/Operadores/HealthWorkerService.cs @@ -324,7 +324,7 @@ namespace AgroBase.Services.Operadores { percent_vel_max = pControle.MovVelocidadeSErvasPercent, percent_vel_min = pControle.MovVelocidadeCErvasPercent, - percent_vel_curva = 20.0, + percent_vel_curva = 15.0, rampa_partida_pct = VariaveisEquipamento.PercentualVelMin, rampa_aceleracao_pct_s = 12.0, rampa_desaceleracao_pct_s = 18.0, @@ -548,7 +548,7 @@ namespace AgroBase.Services.Operadores $"[MAPSYNC] Tentativa={tentativa} | " + $"cmd={commandId} | " + $"rev={mapRevision} | " + - $"pontos={op.Trajetoria._TrajetoriaFixa.Count}" + $"pontos={op.Trajetoria._TrajetoriaFixa?.Count}" ); bool atualizado = diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/camera_worker/camera_oak.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/camera_worker/camera_oak.py index 4cdb02b73..8ad2dfdbb 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/camera_worker/camera_oak.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/camera_worker/camera_oak.py @@ -419,6 +419,7 @@ class CameraOak: t0 = time.time() frame = pkt.getCvFrame() dur = time.time() - t0 + frame_ts_host = time.time() resultado = { "erro": None, @@ -450,7 +451,7 @@ class CameraOak: self.perf.tick( "camera_rgb", latencia_ms=dur * 1000.0, - frame_ts_host=self._rgb_cache_ts, + frame_ts_host=frame_ts_host, frame_ts_device=pkt_ts_device, seq=pkt_seq, seq_delta=seq_delta, @@ -461,7 +462,7 @@ class CameraOak: with self._cache_lock: self._rgb_cache = frame - self._rgb_cache_ts = time.time() + self._rgb_cache_ts = frame_ts_host self._rgb_cache_resultado = resultado self.timestamp_ultimo_frame_rgb = self._rgb_cache_ts self._rgb_seq = pkt_seq @@ -478,6 +479,7 @@ class CameraOak: t0 = time.time() frame = pkt.getFrame() dur = time.time() - t0 + frame_ts_host = time.time() resultado = { "erro": None, @@ -509,7 +511,7 @@ class CameraOak: self.perf.tick( "camera_depth", latencia_ms=dur * 1000.0, - frame_ts_host=self._depth_cache_ts, + frame_ts_host=frame_ts_host, frame_ts_device=pkt_ts_device, seq=pkt_seq, seq_delta=seq_delta, @@ -520,7 +522,7 @@ class CameraOak: with self._cache_lock: self._depth_cache = frame - self._depth_cache_ts = time.time() + self._depth_cache_ts = frame_ts_host self._depth_cache_resultado = resultado self.timestamp_ultimo_frame_depth = self._depth_cache_ts self.ultimo_resultado_depth = resultado @@ -1062,6 +1064,50 @@ class CameraOak: self._definir_heartbeat() return frame, resultado + def requisitar_frame_rgb_ref(self): + """ + Retorna referência READ-ONLY ao latest frame RGB cacheado. + + Diferente de requisitar_frame_rgb(), não copia ~6 MB por frame 1080p. + O cache da câmera apenas substitui a referência por um ndarray novo; não + modifica o ndarray anterior. Portanto o consumidor pode manter a referência + durante resize/inferência sem segurar o lock. + + REGRA: o consumidor não deve escrever no ndarray retornado. + """ + with self._cache_lock: + resultado = dict(self._rgb_cache_resultado) + erro = resultado.get("erro") if self._is_erro_fatal_depthai(resultado.get("erro")) else None + frame = self._rgb_cache + + if erro: + self._falha_fatal_depthai(erro) + return None, resultado + + self._definir_heartbeat() + return frame, resultado + + def requisitar_frame_depth_ref(self): + """Mesmo contrato read-only do RGB, aplicado ao depth latest-frame.""" + if not self.tem_depth: + return None, { + "erro": "Camera nao possui sensor de profundidade", + "duracao": 0, + "frame_valido": False, + } + + with self._cache_lock: + resultado = dict(self._depth_cache_resultado) + erro = resultado.get("erro") if self._is_erro_fatal_depthai(resultado.get("erro")) else None + frame = self._depth_cache + + if erro: + self._falha_fatal_depthai(erro) + return None, resultado + + self._definir_heartbeat() + return frame, resultado + def requisitar_frame_depth(self): if not self.tem_depth: return None, { diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/manager_worker/modulos/movimentacao.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/manager_worker/modulos/movimentacao.py index f5719cf0e..9486e6031 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/manager_worker/modulos/movimentacao.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/manager_worker/modulos/movimentacao.py @@ -1,6 +1,9 @@ import time -from shared.enums import ModoOperacao, StatusModulo, StatusCarroMapa, T_Code +from shared.enums import ( + ModoOperacao, StatusModulo, StatusCarroMapa, T_Code, + TipoMovimentoDirecional, +) from shared.contexto_global_redis import ContextoGlobalRedis, CtxKey from manager_worker.config import mostrar_log from manager_worker.filtros import FiltroVelocidade @@ -67,6 +70,44 @@ def definir_comando(filtro_vel: FiltroVelocidade, erro_orientacao: float = None) abs(_float(dir_cfg.get("angulo_max", 30.0), 30.0)) ) + # V5 - velocidade coerente com a política geométrica do MPC. + # Até ~5° o carro pode manter velocidade de cruzeiro. Entre 5° e 15° + # reduz progressivamente; a partir de 15° (faixa de MovimentoArco) usa + # velocidade de curva. Os parâmetros podem ser publicados em Mov sem + # exigir mudança imediata do contrato C#. + erro_heading_curva_forte_graus = _clamp( + _float( + mov_cfg.get("erro_heading_curva_forte_graus", 15.0), + 15.0, + ), + 3.0, + ang_max, + ) + erro_heading_curva_deadband_graus = _clamp( + _float( + mov_cfg.get("erro_heading_curva_deadband_graus", 5.0), + 5.0, + ), + 0.0, + max(0.0, erro_heading_curva_forte_graus - 0.5), + ) + + # Movimento/ângulo atualmente aplicados. Como o Processador calcula MOV + # antes de DIR, isto funciona como memória segura do ciclo anterior. O + # erro geométrico abaixo continua sendo o gatilho primário do ciclo atual. + tipo_direcional_atual = _enum_or( + TipoMovimentoDirecional, + controle.get( + "tipo_movimento_direcional", + TipoMovimentoDirecional.RodasDianteiras.value, + ), + TipoMovimentoDirecional.RodasDianteiras, + ) + angulo_direcional_atual = abs(_float( + controle.get("angulo_sp", 0.0), + 0.0, + )) + pontos_min_fim_corredor = max( 0, _int(mov_cfg.get("pontos_fim_corredor_reduzir", 3), 3) @@ -105,12 +146,43 @@ def definir_comando(filtro_vel: FiltroVelocidade, erro_orientacao: float = None) status_carro = estado["status_carro"] erro_angular = estado["erro_orientacao"] + erro_angular_caminho = float( + estado.get("erro_angular_caminho", erro_angular) + ) + erro_angular_combinado = float( + estado.get("erro_angular_combinado", erro_angular) + ) pontos_fim_corredor = estado["pontos_fim_corredor"] estado_valido = estado["valido"] debug["status_carro"] = status_carro.name debug["erro_orientacao"] = erro_angular + debug["erro_orientacao_fonte"] = estado.get( + "erro_orientacao_fonte", + "nao_informada" + ) + debug["erros_orientacao"] = { + "selecionado": round(float(erro_angular), 4), + "combinado": round( + float(estado.get("erro_angular_combinado", erro_angular)), + 4 + ), + "caminho": round( + float(estado.get("erro_angular_caminho", erro_angular)), + 4 + ), + "proximo_ponto": round( + float(estado.get("erro_angular_proximo_ponto", erro_angular)), + 4 + ), + } debug["pontos_fim_corredor"] = pontos_fim_corredor + debug["geometria_curva"] = { + "erro_heading_curva_forte_graus": round(erro_heading_curva_forte_graus, 3), + "erro_heading_curva_deadband_graus": round(erro_heading_curva_deadband_graus, 3), + "tipo_direcional_atual": tipo_direcional_atual.name, + "angulo_direcional_atual": round(angulo_direcional_atual, 3), + } if not estado_valido: motivo = estado.get("motivo", "Estado operacional inválido para movimento") @@ -133,10 +205,45 @@ def definir_comando(filtro_vel: FiltroVelocidade, erro_orientacao: float = None) resetar_filtro = False reduzir_para_pulverizar = False + reduzir_por_weed_indisponivel = False + limitar_velocidade_por_weed = False reduzir_para_fim_corredor = False reduzir_por_curva = False + teto_rigido_curva = False + teto_curva_pct = None reduzir_por_ipb = False + # Erro geométrico que conversa com o seletor do MPC. Em tracking de + # caminho e no retorno usamos o heading da centerline, evitando que um + # simples offset lateral pareça uma curva. Em Direcionando preservamos + # o erro combinado porque ali a aquisição do alvo ainda é relevante. + if status_carro in [ + StatusCarroMapa.CaminhandoRua, + StatusCarroMapa.Direcionando, + StatusCarroMapa.RetornandoBase, + ]: + # Mesma referência geométrica do MPC V5: heading da centerline. + # Assim, rover paralelo porém deslocado não é confundido com curva. + erro_curva_graus = abs(erro_angular_caminho) + fonte_curva = "erro_angular_caminho" + else: + erro_curva_graus = abs(erro_angular_combinado) + fonte_curva = "erro_angular_combinado" + + arco_ja_aplicado = ( + tipo_direcional_atual == TipoMovimentoDirecional.MovimentoArco + ) + curva_forte_geometrica = ( + erro_curva_graus >= erro_heading_curva_forte_graus + ) + + debug["geometria_curva"].update({ + "erro_usado_graus": round(float(erro_curva_graus), 3), + "fonte_erro": fonte_curva, + "arco_ja_aplicado": bool(arco_ja_aplicado), + "curva_forte_geometrica": bool(curva_forte_geometrica), + }) + if status_carro == StatusCarroMapa.Parado: velocidade_sp = 0.0 hard_stop = True @@ -151,6 +258,8 @@ def definir_comando(filtro_vel: FiltroVelocidade, erro_orientacao: float = None) ]: velocidade_sp = vel_curva reduzir_por_curva = True + teto_rigido_curva = True + teto_curva_pct = vel_curva debug["motivos"].append("movimento em curva/manobra") elif status_carro in [ @@ -158,26 +267,64 @@ def definir_comando(filtro_vel: FiltroVelocidade, erro_orientacao: float = None) StatusCarroMapa.RetornandoBase ]: velocidade_sp = calcular_velocidade_relativa( - vel_min=vel_com_ervas, + vel_min=vel_curva, vel_max=vel_sem_ervas, - ang_max=ang_max, - erro_orientacao=erro_angular, - k=0.85, + ang_max=erro_heading_curva_forte_graus, + erro_orientacao=erro_curva_graus, + k=1.0, curva=0.80, - dead=10.0 + dead=erro_heading_curva_deadband_graus + ) + reduzir_por_curva = bool( + erro_curva_graus > erro_heading_curva_deadband_graus + or arco_ja_aplicado + ) + teto_rigido_curva = bool(reduzir_por_curva) + if teto_rigido_curva: + teto_curva_pct = ( + vel_curva + if (curva_forte_geometrica or arco_ja_aplicado) + else velocidade_sp + ) + debug["motivos"].append( + f"{status_carro.name}: velocidade pela geometria do caminho/alvo" ) - debug["motivos"].append("direcionando/retornando com redução por erro angular") elif status_carro == StatusCarroMapa.CaminhandoRua: info_ervas = _ler_weed_worker() - ervas_no_radar = ( - pulverizador_automatico and - info_ervas["atualizado"] and - info_ervas["ervas_no_radar"] - ) + # O WeedWorker é a autoridade para dizer se há ervas no radar. + # + # Contrato atual: + # DadosWeedWorker.analise = envelope + # envelope.analise = dados_visuais reais + # + # _ler_weed_worker() normaliza esse contrato e também aceita o + # formato legado (campos diretamente em DadosWeedWorker.analise). + # + # Se o campo booleano não existir em uma versão legada, usamos + # ema_global apenas como fallback. + if info_ervas["campo_ervas_no_radar_presente"]: + ervas_detectadas = bool(info_ervas["ervas_no_radar"]) + else: + ervas_detectadas = ( + info_ervas["percentual_ervas"] >= ervas_min_percent + ) - reduzir_para_pulverizar = bool(ervas_no_radar) + if pulverizador_automatico: + if info_ervas["atualizado"]: + reduzir_para_pulverizar = bool(ervas_detectadas) + else: + # Falha conservadora de qualidade de aplicação: + # percepção desconhecida NÃO significa "sem ervas". + # Mantém a velocidade de pulverização até a análise voltar, + # sem parar a operação por si só. + reduzir_por_weed_indisponivel = True + + limitar_velocidade_por_weed = bool( + reduzir_para_pulverizar + or reduzir_por_weed_indisponivel + ) reduzir_para_fim_corredor = ( pontos_fim_corredor >= 0 and @@ -185,20 +332,57 @@ def definir_comando(filtro_vel: FiltroVelocidade, erro_orientacao: float = None) ) debug["weed"] = info_ervas + debug["weed"]["pulverizador_automatico"] = bool( + pulverizador_automatico + ) + debug["weed"]["ervas_detectadas_efetivo"] = bool( + ervas_detectadas + ) + debug["weed"]["limitar_velocidade"] = bool( + limitar_velocidade_por_weed + ) velocidade_livre = calcular_velocidade_relativa( - vel_min=vel_com_ervas, + vel_min=vel_curva, vel_max=vel_sem_ervas, - ang_max=ang_max, - erro_orientacao=erro_angular, + ang_max=erro_heading_curva_forte_graus, + erro_orientacao=erro_curva_graus, k=1.0, - curva=0.70, - dead=5.0 + curva=0.80, + dead=erro_heading_curva_deadband_graus + ) + + # Mesmo que o C# publique CaminhandoRua durante RetornoBase, uma + # curva forte da centerline não pode ser tratada como reta rápida. + # O heading do caminho é a mesma referência usada pelo MPC V5. + reduzir_por_curva = bool( + erro_curva_graus > erro_heading_curva_deadband_graus + or arco_ja_aplicado + ) + teto_rigido_curva = bool(reduzir_por_curva) + if teto_rigido_curva: + teto_curva_pct = ( + vel_curva + if (curva_forte_geometrica or arco_ja_aplicado) + else velocidade_livre + ) + + debug["motivos"].append( + "caminhando rua: velocidade por erro angular do caminho" ) if reduzir_para_pulverizar: velocidade_sp = min(velocidade_livre, vel_com_ervas) - debug["motivos"].append("ervas no radar: reduzindo para pulverizar") + debug["motivos"].append( + "ervas no radar: reduzindo para pulverizar" + ) + + elif reduzir_por_weed_indisponivel: + velocidade_sp = min(velocidade_livre, vel_com_ervas) + debug["motivos"].append( + "WeedWorker ausente/desatualizado: " + "mantendo velocidade conservadora de pulverização" + ) elif reduzir_para_fim_corredor: velocidade_sp = min(velocidade_livre, vel_com_ervas) @@ -217,6 +401,16 @@ def definir_comando(filtro_vel: FiltroVelocidade, erro_orientacao: float = None) f"status_carro não tratado com segurança: {status_carro.name}" ) + debug["geometria_curva"].update({ + "reduzir_por_curva": bool(reduzir_por_curva), + "teto_rigido_curva": bool(teto_rigido_curva), + "teto_curva_pct": ( + None if teto_curva_pct is None + else round(float(teto_curva_pct), 3) + ), + "vel_curva_pct": round(float(vel_curva), 3), + }) + # Se já decidiu parar, não faz limitador tentar reviver velocidade. if hard_stop: _resetar_rampa_velocidade(filtro_vel) @@ -284,7 +478,8 @@ def definir_comando(filtro_vel: FiltroVelocidade, erro_orientacao: float = None) if usar_imu_movimento: vel_imu, imu_resetar, imu_stop, imu_debug = _aplicar_limitador_imu( - velocidade_atual=velocidade_sp + velocidade_atual=velocidade_sp, + vel_min_equip=vel_min_equip ) debug["imu"] = imu_debug @@ -329,7 +524,7 @@ def definir_comando(filtro_vel: FiltroVelocidade, erro_orientacao: float = None) # Aceleração normal é filtrada para não dar tranco. reducao_seguranca = bool( - reduzir_para_pulverizar + limitar_velocidade_por_weed or reduzir_para_fim_corredor or reduzir_por_curva or reduzir_por_ipb @@ -351,23 +546,59 @@ def definir_comando(filtro_vel: FiltroVelocidade, erro_orientacao: float = None) debug=debug, ) - # Correção crítica: - # havendo erva no radar, a saída nunca pode superar a mínima definida. - if reduzir_para_pulverizar: - velocidade_filtrada = min( - velocidade_filtrada, - vel_com_ervas - ) + # Teto rígido de curva/Arco. A rampa de segurança continua suavizando + # reduções moderadas, mas uma geometria já classificada como curva forte + # não pode passar um ciclo em velocidade de cruzeiro enquanto o MPC + # coloca o rover em MovimentoArco. + if teto_rigido_curva and teto_curva_pct is not None: + teto_curva_pct = _clamp(teto_curva_pct, vel_min_equip, vel_sem_ervas) + velocidade_filtrada = min(velocidade_filtrada, teto_curva_pct) - # Sincroniza o estado da rampa com a saída realmente aplicada. try: filtro_vel._ultimo_sp_campo = velocidade_filtrada filtro_vel._ultimo_ts_campo = time.monotonic() except Exception: pass + motivo_curva = ( + "curva forte/Arco" + if (curva_forte_geometrica or arco_ja_aplicado) + else "curva moderada" + ) debug["limitadores"].append( - f"ervas no radar: teto rígido em {vel_com_ervas:.1f}%" + f"{motivo_curva}: teto geométrico em {teto_curva_pct:.1f}%" + ) + + # Teto rígido do WeedWorker: + # + # - ervas detectadas: nunca ultrapassa vel_com_ervas; + # - WeedWorker desatualizado com pulverizador automático ligado: + # também não assume que a rua está livre. + # + # Esse clamp acontece DEPOIS da rampa para não existir overshoot + # transitório acima da velocidade de aplicação. + if limitar_velocidade_por_weed: + velocidade_filtrada = min( + velocidade_filtrada, + vel_com_ervas + ) + + # Sincroniza o estado da rampa com a saída realmente aplicada. + # Assim, quando a condição desaparecer, a aceleração volta a + # acontecer pela rampa normal em vez de saltar para o SP antigo. + try: + filtro_vel._ultimo_sp_campo = velocidade_filtrada + filtro_vel._ultimo_ts_campo = time.monotonic() + except Exception: + pass + + if reduzir_para_pulverizar: + motivo_teto = "ervas no radar" + else: + motivo_teto = "WeedWorker indisponível/desatualizado" + + debug["limitadores"].append( + f"{motivo_teto}: teto rígido em {vel_com_ervas:.1f}%" ) velocidade_filtrada = round( @@ -415,6 +646,33 @@ def _resolver_estado_operacional( contexto, erro_orientacao ): + """ + Resolve o estado operacional e escolhe a referência angular usada + EXCLUSIVAMENTE para a política de velocidade. + + Política alinhada ao MPC Path Tracking V2: + + - CaminhandoRua: + usa erro_angular_caminho. + O que importa para liberar velocidade é o corpo do rover estar + alinhado com a orientação local da passada. Deslocamento lateral + pode ser corrigido pelo MovimentoDiagonal sem penalizar a velocidade + como se o rover estivesse "torto". + + - Direcionando / RetornandoBase: + usa erro_angular combinado. + Fora da passada, a aquisição geométrica do alvo continua relevante. + + - EntrandoRua / SaindoRua / Manobrando: + a velocidade base é vel_curva, portanto o erro selecionado não + altera a velocidade nesses estados. Mantemos o combinado para + diagnóstico/coerência. + + - erro_orientacao explícito: + continua tendo prioridade para preservar compatibilidade com + chamadas externas/legadas. + """ + if op_modo in [ModoOperacao.MapaGPS, ModoOperacao.RetornoBase]: trajetoria = _dict(contexto.get("Trajetoria", {})) @@ -424,6 +682,10 @@ def _resolver_estado_operacional( "motivo": "Trajetoria ausente no contexto", "status_carro": StatusCarroMapa.Parado, "erro_orientacao": 0.0, + "erro_orientacao_fonte": "trajetoria_ausente", + "erro_angular_combinado": 0.0, + "erro_angular_caminho": 0.0, + "erro_angular_proximo_ponto": 0.0, "pontos_fim_corredor": -1, } @@ -433,33 +695,86 @@ def _resolver_estado_operacional( StatusCarroMapa.Parado ) - if erro_orientacao is None: - erro_orientacao = _float(trajetoria.get("erro_angular", 0.0), 0.0) + erro_combinado = _float( + trajetoria.get("erro_angular", 0.0), + 0.0 + ) + erro_caminho = _float( + trajetoria.get( + "erro_angular_caminho", + erro_combinado + ), + erro_combinado + ) + erro_proximo_ponto = _float( + trajetoria.get( + "erro_angular_proximo_ponto", + erro_combinado + ), + erro_combinado + ) + + if erro_orientacao is not None: + erro_selecionado = _float( + erro_orientacao, + erro_combinado + ) + fonte_erro = "override_explicito" + + elif status_carro == StatusCarroMapa.CaminhandoRua: + # Na passada, o heading local da rua é a referência correta. + # O bearing para waypoint não deve reduzir velocidade só porque + # o MPC está corrigindo offset lateral em MovimentoDiagonal. + erro_selecionado = erro_caminho + fonte_erro = "erro_angular_caminho" + + else: + # Aquisição, retorno e demais estados mantêm a referência + # combinada que incorpora a geometria do alvo. + erro_selecionado = erro_combinado + fonte_erro = "erro_angular_combinado" corredor = _dict(trajetoria.get("CorredorAtual", {})) - pontos_fim_corredor = _int(corredor.get("pontos_restantes", -1), -1) + pontos_fim_corredor = _int( + corredor.get("pontos_restantes", -1), + -1 + ) return { "valido": True, "motivo": "", "status_carro": status_carro, - "erro_orientacao": _float(erro_orientacao, 0.0), + "erro_orientacao": float(erro_selecionado), + "erro_orientacao_fonte": fonte_erro, + "erro_angular_combinado": float(erro_combinado), + "erro_angular_caminho": float(erro_caminho), + "erro_angular_proximo_ponto": float(erro_proximo_ponto), "pontos_fim_corredor": pontos_fim_corredor, } if op_modo == ModoOperacao.MapeamentoVisual: - dados_vw = _dict(ContextoGlobalRedis.get(CtxKey.DadosVisualWorker, {})) - segmentacao = _dict(dados_vw.get("segmentacao", {})) + dados_vw = _dict( + ContextoGlobalRedis.get(CtxKey.DadosVisualWorker, {}) + ) + segmentacao = _dict( + dados_vw.get("segmentacao", {}) + ) seg_ts = _float( segmentacao.get( "ts", - dados_vw.get("ts_segmentacao", dados_vw.get("momento", 0.0)) + dados_vw.get( + "ts_segmentacao", + dados_vw.get("momento", 0.0) + ) ), 0.0 ) - seg_atualizada, idade_ms = _timestamp_atualizado(seg_ts, max_idade_s=1.0) + seg_atualizada, idade_ms = _timestamp_atualizado( + seg_ts, + max_idade_s=1.0 + ) if not segmentacao or not seg_atualizada: motivo = "segmentacao visual ausente/desatualizada" @@ -472,31 +787,61 @@ def _resolver_estado_operacional( "motivo": motivo, "status_carro": StatusCarroMapa.Parado, "erro_orientacao": 0.0, + "erro_orientacao_fonte": "segmentacao_invalida", + "erro_angular_combinado": 0.0, + "erro_angular_caminho": 0.0, + "erro_angular_proximo_ponto": 0.0, "pontos_fim_corredor": -1, } status_carro = _enum_or( StatusCarroMapa, - segmentacao.get("status_corredor", StatusCarroMapa.Parado.value), + segmentacao.get( + "status_corredor", + StatusCarroMapa.Parado.value + ), StatusCarroMapa.Parado ) - if erro_orientacao is None: - erro_orientacao = _float(segmentacao.get("erro_angular", 0.0), 0.0) + erro_visual = _float( + segmentacao.get("erro_angular", 0.0), + 0.0 + ) + + if erro_orientacao is not None: + erro_selecionado = _float( + erro_orientacao, + erro_visual + ) + fonte_erro = "override_explicito" + else: + erro_selecionado = erro_visual + fonte_erro = "segmentacao_visual" return { "valido": True, "motivo": "", "status_carro": status_carro, - "erro_orientacao": _float(erro_orientacao, 0.0), + "erro_orientacao": float(erro_selecionado), + "erro_orientacao_fonte": fonte_erro, + "erro_angular_combinado": float(erro_visual), + "erro_angular_caminho": float(erro_visual), + "erro_angular_proximo_ponto": float(erro_visual), "pontos_fim_corredor": -1, } return { "valido": False, - "motivo": f"Modo de operação sem estratégia de movimento: {op_modo.name}", + "motivo": ( + "Modo de operação sem estratégia de movimento: " + f"{op_modo.name}" + ), "status_carro": StatusCarroMapa.Parado, "erro_orientacao": 0.0, + "erro_orientacao_fonte": "modo_sem_estrategia", + "erro_angular_combinado": 0.0, + "erro_angular_caminho": 0.0, + "erro_angular_proximo_ponto": 0.0, "pontos_fim_corredor": -1, } @@ -506,34 +851,119 @@ def _resolver_estado_operacional( # ============================================================ def _ler_weed_worker(): - try: - dados = _dict(ContextoGlobalRedis.get(CtxKey.DadosWeedWorker, {})) - analise = _dict(dados.get("analise", {})) + """ + Lê e normaliza o contrato do WeedWorker. + Contrato atual publicado pelo CameraManager: + + DadosWeedWorker = { + "ts_publicacao": ..., + "ts_analise": ..., + "analise": { + "ts_analise": ..., + "fps_model": ..., + "fps_inferencia": ..., + "infer_ms": ..., + "infer_gpu_ms": ..., + "analise": { + "ervas_no_radar": ..., + "estatisticas": { + "erva": { + "ema_global": ... + } + }, + ... + } + } + } + + Também aceita o formato legado, no qual ervas_no_radar e estatisticas + ficam diretamente em DadosWeedWorker["analise"]. + """ + try: + dados = _dict( + ContextoGlobalRedis.get(CtxKey.DadosWeedWorker, {}) + ) + + envelope = _dict(dados.get("analise", {})) + + # Contrato atual possui envelope["analise"]. + # No contrato legado, o próprio envelope já é a análise. + analise_interna = _dict(envelope.get("analise", {})) + + if analise_interna: + analise = analise_interna + esquema = "nested_v2" + else: + analise = envelope + esquema = "flat_legacy" + + # Preferência de timestamp: + # 1) timestamp do envelope da análise atual; + # 2) timestamp dentro da análise legada; + # 3) timestamp raiz do DadosWeedWorker. ts = _float( - analise.get( - "ts", - analise.get( - "timestamp", - dados.get("ts_analise", dados.get("momento", 0.0)) + envelope.get( + "ts_analise", + envelope.get( + "ts", + envelope.get( + "timestamp", + analise.get( + "ts", + analise.get( + "timestamp", + dados.get( + "ts_analise", + dados.get( + "ts_publicacao", + dados.get("momento", 0.0) + ) + ) + ) + ) + ) ) ), 0.0 ) - atualizado, idade_ms = _timestamp_atualizado(ts, max_idade_s=1.5) + atualizado, idade_ms = _timestamp_atualizado( + ts, + max_idade_s=1.5 + ) - estatisticas = _dict(analise.get("estatisticas", {})) - erva = _dict(estatisticas.get("erva", {})) + estatisticas = _dict( + analise.get("estatisticas", {}) + ) + erva = _dict( + estatisticas.get("erva", {}) + ) - ervas_no_radar = _bool(analise.get("ervas_no_radar", False)) - percentual_ervas = _float(erva.get("ema_global", 0.0), 0.0) + campo_ervas_presente = "ervas_no_radar" in analise + + ervas_no_radar = _bool( + analise.get("ervas_no_radar", False) + ) + + percentual_ervas = _float( + erva.get( + "ema_global", + analise.get("percentual_ervas_no_radar", 0.0) + ), + 0.0 + ) return { - "atualizado": atualizado, + "atualizado": bool(atualizado), "idade_ms": idade_ms, - "ervas_no_radar": ervas_no_radar, + "ervas_no_radar": bool(ervas_no_radar), + "campo_ervas_no_radar_presente": bool( + campo_ervas_presente + ), "percentual_ervas": percentual_ervas, + "timestamp_analise": ts, + "esquema": esquema, } except Exception as e: @@ -541,7 +971,10 @@ def _ler_weed_worker(): "atualizado": False, "idade_ms": None, "ervas_no_radar": False, + "campo_ervas_no_radar_presente": False, "percentual_ervas": 0.0, + "timestamp_analise": 0.0, + "esquema": "erro", "erro": str(e), } @@ -643,7 +1076,7 @@ def _aplicar_limitador_visual( # Limitador IMU # ============================================================ -def _aplicar_limitador_imu(*, velocidade_atual): +def _aplicar_limitador_imu(*, velocidade_atual, vel_min_equip): debug = { "habilitado": True, "aplicado": False, @@ -693,7 +1126,41 @@ def _aplicar_limitador_imu(*, velocidade_atual): if risk_level_num <= 0: return velocidade_atual, False, False, debug - vel_nova = float(velocidade_atual) * velocidade_factor + velocidade_atual = max(0.0, float(velocidade_atual)) + vel_min_equip = max(0.0, float(vel_min_equip)) + + # Attention/risk são REDUÇÕES, não ordens de parada. + # + # Se a IMU, sozinha, levar o SP abaixo da velocidade mínima útil do + # equipamento, limita no piso operacional em vez de transformar um + # nível attention/risk em parada. + # + # O min(..., velocidade_atual) é proposital: a IMU nunca pode elevar + # uma velocidade que já tenha sido reduzida por outro limitador + # (Visual Worker, IPB, fim de corredor, pulverização etc.). + vel_calculada = velocidade_atual * velocidade_factor + + if velocidade_atual >= vel_min_equip and vel_calculada > 0.0: + vel_nova = min( + velocidade_atual, + max(vel_min_equip, vel_calculada) + ) + else: + # Se outro limitador já entregou algo abaixo do piso, a IMU não + # "ressuscita" a velocidade. A regra geral posterior decide se + # esse valor deve virar parada. + vel_nova = min(velocidade_atual, vel_calculada) + + debug["velocidade_entrada"] = round(velocidade_atual, 3) + debug["velocidade_calculada"] = round(vel_calculada, 3) + debug["velocidade_saida"] = round(vel_nova, 3) + debug["vel_min_equip"] = round(vel_min_equip, 3) + debug["piso_operacional_aplicado"] = bool( + velocidade_atual >= vel_min_equip + and 0.0 < vel_calculada < vel_min_equip + and abs(vel_nova - vel_min_equip) < 1e-9 + ) + resetar_filtro = False return vel_nova, resetar_filtro, False, debug diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/manager_worker/modulos/mpc.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/manager_worker/modulos/mpc.py index 8ccf37476..ee460033f 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/manager_worker/modulos/mpc.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/manager_worker/modulos/mpc.py @@ -341,7 +341,156 @@ class ControladorMPC: max_value=20.0, ) - # Amortecimento aplicado somente durante CaminhandoRua. Nas curvas o + # ------------------------------------------------------------------ + # PATH TRACKING V2 - referência geométrica da passada + # ------------------------------------------------------------------ + # Na passada o rover NÃO mira um ponto. O ponto/progresso continua + # sendo autoridade do C#, mas a direção segue: (1) heading médio do + # caminho nos próximos metros + (2) cross-track lateral. O alvo + # pure-pursuit fica reservado para entrada/saída/manobra. + self._heading_lookahead_base_m = _safe_float( + parametros_mpc.get("heading_lookahead_base_m", 1.60), + 1.60, min_value=0.40, max_value=8.0, + ) + self._heading_lookahead_velocidade_s = _safe_float( + parametros_mpc.get("heading_lookahead_velocidade_s", 0.55), + 0.55, min_value=0.0, max_value=4.0, + ) + self._heading_lookahead_min_m = _safe_float( + parametros_mpc.get("heading_lookahead_min_m", 1.20), + 1.20, min_value=0.25, max_value=8.0, + ) + self._heading_lookahead_max_m = _safe_float( + parametros_mpc.get("heading_lookahead_max_m", 3.00), + 3.00, min_value=self._heading_lookahead_min_m, max_value=12.0, + ) + + # Heading primeiro, lateral depois. A diagonal só pode existir quando + # o corpo já está praticamente paralelo ao caminho; assim ela + # recentraliza sem congelar um erro de heading de 5-10 graus. + self._heading_deadband_graus = _safe_float( + parametros_mpc.get("heading_deadband_graus", 0.45), + 0.45, min_value=0.0, max_value=3.0, + ) + self._diagonal_entrada_heading_graus = _safe_float( + parametros_mpc.get("diagonal_entrada_heading_graus", 1.25), + 1.25, min_value=0.2, max_value=8.0, + ) + self._diagonal_saida_heading_graus = _safe_float( + parametros_mpc.get("diagonal_saida_heading_graus", 2.50), + 2.50, min_value=self._diagonal_entrada_heading_graus, max_value=12.0, + ) + self._lateral_deadband_m = _safe_float( + parametros_mpc.get("lateral_deadband_m", 0.05), + 0.05, min_value=0.0, max_value=0.30, + ) + self._ganho_heading_reta = _safe_float( + parametros_mpc.get("ganho_heading_reta", 1.00), + 1.00, min_value=0.1, max_value=3.0, + ) + self._ganho_lateral_dianteira = _safe_float( + parametros_mpc.get("ganho_lateral_dianteira", 0.32), + 0.32, min_value=0.0, max_value=3.0, + ) + self._ganho_lateral_diagonal = _safe_float( + parametros_mpc.get("ganho_lateral_diagonal", 0.85), + 0.85, min_value=0.0, max_value=4.0, + ) + self._velocidade_offset_lateral = _safe_float( + parametros_mpc.get("velocidade_offset_lateral", 0.35), + 0.35, min_value=0.05, max_value=2.0, + ) + self._angulo_diagonal_max_graus = _safe_float( + parametros_mpc.get("angulo_diagonal_max_graus", 16.0), + 16.0, min_value=1.0, max_value=self.angulo_max_graus, + ) + self._heading_reta_aquisicao_max_graus = _safe_float( + parametros_mpc.get("heading_reta_aquisicao_max_graus", 15.0), + 15.0, min_value=3.0, max_value=35.0, + ) + self._lateral_reta_aquisicao_max_m = _safe_float( + parametros_mpc.get("lateral_reta_aquisicao_max_m", 0.80), + 0.80, min_value=0.15, max_value=2.0, + ) + self._penalidade_troca_tipo = _safe_float( + parametros_mpc.get("penalidade_troca_tipo", 0.18), + 0.18, min_value=0.0, max_value=3.0, + ) + + # Direcionando fora da rua: usa Arco para a correção grosseira de + # heading e volta para RodasDianteiras somente quando já estiver bem + # alinhado. Os dois limiares criam histerese e evitam troca de modo + # perto da fronteira. + self._direcionando_arco_entrada_graus = _safe_float( + parametros_mpc.get("direcionando_arco_entrada_graus", 15.0), + 15.0, min_value=3.0, max_value=60.0, + ) + self._direcionando_arco_saida_graus = _safe_float( + parametros_mpc.get("direcionando_arco_saida_graus", 8.0), + 8.0, min_value=1.0, max_value=self._direcionando_arco_entrada_graus, + ) + + # -------------------------------------------------------------- + # V5 - Política geométrica universal do path tracking. + # + # A lógica vale tanto para a passada agrícola quanto para a rota geral + # de RetornoBase: + # - heading muito errado -> MovimentoArco; + # - heading moderadamente errado -> RodasDianteiras; + # - heading praticamente alinhado + cross-track -> MovimentoDiagonal. + # + # A histerese 15°/8° evita ficar trocando Arco <-> Dianteira na + # fronteira. Os defaults herdam a política já validada em Direcionando. + # Manobras estruturais reais continuam com Arco obrigatório. + # -------------------------------------------------------------- + self._tracking_arco_entrada_graus = _safe_float( + parametros_mpc.get( + "tracking_arco_entrada_graus", + self._direcionando_arco_entrada_graus, + ), + self._direcionando_arco_entrada_graus, + min_value=3.0, + max_value=60.0, + ) + self._tracking_arco_saida_graus = _safe_float( + parametros_mpc.get( + "tracking_arco_saida_graus", + self._direcionando_arco_saida_graus, + ), + self._direcionando_arco_saida_graus, + min_value=1.0, + max_value=self._tracking_arco_entrada_graus, + ) + + # -------------------------------------------------------------- + # V4 - Entrada/saída de rua: Arco + RodasDianteiras competem. + # + # A preferência usa o HEADING DO CAMINHO, e não o bearing até o + # alvo. Assim, quando o corpo já está alinhado com a próxima rua, + # RodasDianteiras ganham vantagem e terminam a aquisição de forma + # suave. Arco continua favorecido quando o heading ainda está muito + # errado. + # + # Entre os dois limiares a preferência varia continuamente; a + # penalidade normal de troca de tipo fornece histerese adicional. + # -------------------------------------------------------------- + self._transicao_rua_dianteira_heading_graus = _safe_float( + parametros_mpc.get("transicao_rua_dianteira_heading_graus", 7.0), + 7.0, min_value=0.5, max_value=30.0, + ) + self._transicao_rua_arco_heading_graus = _safe_float( + parametros_mpc.get("transicao_rua_arco_heading_graus", 20.0), + 20.0, + min_value=self._transicao_rua_dianteira_heading_graus + 0.5, + max_value=60.0, + ) + self._transicao_rua_penalidade_modo = _safe_float( + parametros_mpc.get("transicao_rua_penalidade_modo", 1.0), + 1.0, min_value=0.0, max_value=5.0, + ) + + # Amortecimento aplicado ao tracking de passada. MovimentoArco em + # manobra preserva autoridade; RodasDianteiras/Diagonal são limitados. # MPC preserva toda a autoridade angular necessaria para a manobra. self._filtro_angulo_reta_tau_s = _safe_float( parametros_mpc.get("filtro_angulo_reta_tau_s", 0.25), @@ -562,6 +711,11 @@ class ControladorMPC: return (float(xy[0]), float(xy[1])), i def _distancia_lookahead(self, status_carro, velocidade): + """Look-ahead do alvo de manobra/pure-pursuit. + + Em CaminhandoRua esse alvo continua existindo para debug/progresso, + mas NÃO define o heading desejado do rover. + """ v = max(0.0, _safe_float(velocidade, 0.0)) if _status_in(status_carro, [StatusCarroMapa.CaminhandoRua]): @@ -580,7 +734,7 @@ class ControladorMPC: )) def _construir_alvo_direcional(self, x, y, idx_base, status_carro, velocidade): - """Cria alvo pure-pursuit continuo sem alterar o progresso visitado.""" + """Cria alvo contínuo para manobra/debug sem alterar progresso visitado.""" n = len(self.pontos_info) if n <= 0: return {"xy": (float(x), float(y)), "_idx_base": 0} @@ -606,21 +760,285 @@ class ControladorMPC: alvo["_lookahead_m"] = float(max(0.0, s_alvo - s_proj)) return alvo + def _ponto_estrutural(self, idx): + if not self.pontos_info: + return True + i = max(0, min(_safe_int(idx, 0), len(self.pontos_info) - 1)) + p = self.pontos_info[i] + return bool(p.get("estrutural", p.get("PontoEstrutural", False))) + + def _idx_corredor(self, idx): + if not self.pontos_info: + return -1 + i = max(0, min(_safe_int(idx, 0), len(self.pontos_info) - 1)) + p = self.pontos_info[i] + return _safe_int(p.get("idxCorredor", p.get("IdxCorredor", -1)), -1) + + def _segmento_operacional_mesmo_corredor(self, idx_base): + """True quando há um segmento de passada válido junto ao índice base. + + Aceita a borda de entrada como início da aquisição se o segmento à + frente já pertence ao mesmo corredor e o próximo ponto é operacional. + Isso evita obrigar Arco só porque o estado C# ainda diz Manobrando. + """ + n = len(self.pontos_info) + if n < 2: + return False + idx = max(0, min(_safe_int(idx_base, 0), n - 1)) + + pares = [] + if idx + 1 < n: + pares.append((idx, idx + 1)) + if idx - 1 >= 0: + pares.append((idx - 1, idx)) + + for a, b in pares: + ca, cb = self._idx_corredor(a), self._idx_corredor(b) + if ca < 0 or cb < 0 or ca != cb: + continue + # Um ponto estrutural isolado na entrada pode iniciar a reta, mas + # dois pontos estruturais seguidos caracterizam ligação/manobra. + if self._ponto_estrutural(a) and self._ponto_estrutural(b): + continue + return True + return False + + def _segmento_rua_puro_mesmo_corredor(self, idx_base): + """ + True somente quando a referência local está realmente em um trecho + operacional Rua -> Rua do MESMO corredor. + + Diferente de _segmento_operacional_mesmo_corredor(), aqui não aceitamos + ponto estrutural isolado. Isso é importante para distinguir: + + - recuperação dentro/ao lado de uma passada reta; + - entrada, saída, ligação ou cabeceira real. + + Na recuperação de uma passada, MovimentoArco é indesejado: ele dobra a + autoridade de yaw e pode transformar um desvio lateral em uma inversão + brusca de heading. + """ + n = len(self.pontos_info) + if n < 2: + return False + + idx = max(0, min(_safe_int(idx_base, 0), n - 1)) + + pares = [] + if idx + 1 < n: + pares.append((idx, idx + 1)) + if idx - 1 >= 0: + pares.append((idx - 1, idx)) + + for a, b in pares: + ca, cb = self._idx_corredor(a), self._idx_corredor(b) + + if ca < 0 or cb < 0 or ca != cb: + continue + + if self._ponto_estrutural(a) or self._ponto_estrutural(b): + continue + + return True + + return False + + def _limitar_s_ao_corredor(self, s_proj, s_desejado, idx_base): + """Não deixa o heading look-ahead enxergar a rua seguinte pela cabeceira.""" + if self._s_nodes is None or len(self._s_nodes) == 0: + return float(s_desejado) + + n = len(self.pontos_info) + idx = max(0, min(_safe_int(idx_base, 0), n - 1)) + corredor = self._idx_corredor(idx) + if corredor < 0: + return float(s_desejado) + + s_lim = float(self._s_nodes[-1]) + for j in range(idx + 1, n): + if self._idx_corredor(j) != corredor or self._ponto_estrutural(j): + s_lim = float(self._s_nodes[j]) + break + + # precisa sobrar um pequeno vetor para formar direção; se já estamos + # no fim, a tangente local será usada pelo chamador. + return float(min(float(s_desejado), s_lim)) + + def _referencia_caminho_lookahead(self, x, y, idx_referencia, velocidade): + """Retorna heading médio dos próximos metros + cross-track. + + A orientação é a corda da centerline entre a projeção atual e um ponto + adiante. Em reta ela é constante; em curva antecipa suavemente o arco. + Nunca aponta do ROBÔ para o alvo, portanto deslocamento lateral não + contamina o erro de heading. + """ + e_lat, s_proj, i_seg, proj, theta_local, margem = self._projetar_na_trajetoria_local( + x, y, idx_referencia=idx_referencia, + ) + + v = max(0.0, _safe_float(velocidade, 0.0)) + look = self._heading_lookahead_base_m + self._heading_lookahead_velocidade_s * v + look = float(np.clip(look, self._heading_lookahead_min_m, self._heading_lookahead_max_m)) + + s_fim = min(float(self._s_nodes[-1]), float(s_proj) + look) + s_fim = self._limitar_s_ao_corredor(s_proj, s_fim, idx_referencia) + + p0, _ = self._xy_na_abscissa(s_proj) + p1, _ = self._xy_na_abscissa(s_fim) + dx = float(p1[0] - p0[0]) + dy = float(p1[1] - p0[1]) + dist = math.hypot(dx, dy) + + if dist >= 0.20: + theta_ref = float(math.atan2(dx, dy)) + look_real = float(s_fim - s_proj) + else: + theta_ref = float(theta_local) + look_real = 0.0 + + return { + "e_lat": float(e_lat), + "s_proj": float(s_proj), + "i_seg": int(i_seg), + "projecao": proj, + "theta_local": float(theta_local), + "theta_ref": float(theta_ref), + "margem": float(margem), + "lookahead_heading_m": float(max(0.0, look_real)), + } + + def _erro_heading_caminho(self, theta_robo, theta_ref): + return float(self._wrap_pi(float(theta_ref) - float(theta_robo))) + + def _tracking_reta_permitido(self, contexto, idx_base, e_lat, erro_heading_rad): + """ + Decide se o path tracking geométrico deve assumir o controle. + + Regras V5: + + 1) CaminhandoRua / EntrandoRua / SaindoRua / Direcionando / + RetornandoBase: + usam a mesma política geométrica universal. O estado operacional + continua útil para velocidade e semântica, mas não deve aprisionar + a direção em uma geometria específica. + + 2) Manobrando em trecho Rua -> Rua do MESMO corredor: + é recuperação de trajetória e também pode usar a política universal. + + 3) Manobra estrutural real: + fica fora deste ramo e preserva MovimentoArco obrigatório, adequado + à curva de cabeceira/180 graus. + """ + carro = _as_dict(_as_dict(contexto).get("Carro", {})) + status = carro.get("Status", StatusCarroMapa.Parado.value) + + if _status_in( + status, + [ + StatusCarroMapa.CaminhandoRua, + StatusCarroMapa.EntrandoRua, + StatusCarroMapa.SaindoRua, + StatusCarroMapa.Direcionando, + StatusCarroMapa.RetornandoBase, + ], + ): + return True + + if _status_in(status, [StatusCarroMapa.Manobrando]): + return self._segmento_rua_puro_mesmo_corredor(idx_base) + + return False + + def _referencia_direcional_reta(self, e_lat, erro_heading_rad, velocidade, tipo_preferido=None): + """Escolhe a geometria pelo estado geométrico rover x caminho. + + Política universal V5: + - heading muito errado -> MovimentoArco para recuperar orientação com + menor raio; + - heading moderadamente errado -> RodasDianteiras; + - heading praticamente alinhado + cross-track -> MovimentoDiagonal; + - heading alinhado e centralizado -> RodasDianteiras praticamente reta. + + A decisão usa heading do CAMINHO, nunca bearing rover->waypoint. + """ + e_lat = float(e_lat) + v = max(0.0, float(velocidade)) + erro_heading_rad = float(erro_heading_rad) + e_head_deg = abs(math.degrees(erro_heading_rad)) + tipo_prev = _movimento_from_value(tipo_preferido) + + # 1) Erro grande de heading tem prioridade absoluta sobre cross-track. + # A histerese mantém Arco até o erro cair bem abaixo do limiar de + # entrada, evitando caça de geometria em torno de 15 graus. + manter_arco = ( + tipo_prev == TipoMovimentoDirecional.MovimentoArco + and e_head_deg > self._tracking_arco_saida_graus + ) + entrar_arco = e_head_deg >= self._tracking_arco_entrada_graus + + if entrar_arco or manter_arco: + delta = self._ganho_heading_reta * erro_heading_rad + lim = math.radians(self.angulo_max_graus) + return ( + TipoMovimentoDirecional.MovimentoArco, + float(np.clip(delta, -lim, lim)), + ) + + # 2) Só usa diagonal quando o corpo já está praticamente paralelo ao + # caminho. Ela corrige deslocamento lateral sem criar yaw desnecessário. + precisa_recentralizar = abs(e_lat) > self._lateral_deadband_m + manter_diag = ( + tipo_prev == TipoMovimentoDirecional.MovimentoDiagonal + and e_head_deg <= self._diagonal_saida_heading_graus + ) + entrar_diag = ( + precisa_recentralizar + and e_head_deg <= self._diagonal_entrada_heading_graus + ) + usar_diag = bool(manter_diag or entrar_diag) + + if usar_diag: + if abs(e_lat) <= self._lateral_deadband_m: + delta = 0.0 + else: + delta = math.atan2( + self._ganho_lateral_diagonal * e_lat, + v + self._velocidade_offset_lateral, + ) + lim = math.radians(min(self._angulo_diagonal_max_graus, self.angulo_max_graus)) + return TipoMovimentoDirecional.MovimentoDiagonal, float(np.clip(delta, -lim, lim)) + + # 3) Faixa intermediária: RodasDianteiras têm UMA responsabilidade, + # corrigir heading. O erro lateral fica reservado para a diagonal quando + # o corpo voltar à janela estreita de alinhamento. + if e_head_deg <= self._heading_deadband_graus: + e_head_eff = 0.0 + else: + e_head_eff = erro_heading_rad + + delta = self._ganho_heading_reta * e_head_eff + lim = math.radians(self.angulo_max_graus) + return TipoMovimentoDirecional.RodasDianteiras, float(np.clip(delta, -lim, lim)) + def _estabilizar_angulo_saida( self, angulo_desejado_rad, status_carro, comando_anterior, dt_controle, + tipo_movimento=None, ): - """Filtra/rate-limita esterco somente na passada reta.""" + """Rate-limit de saída para modos de tracking; Arco mantém autoridade.""" desejado = float(np.clip( angulo_desejado_rad, -np.radians(self.angulo_max_graus), np.radians(self.angulo_max_graus), )) - if not _status_in(status_carro, [StatusCarroMapa.CaminhandoRua]): + tipo = _movimento_from_value(tipo_movimento) if tipo_movimento is not None else None + if tipo == TipoMovimentoDirecional.MovimentoArco: + return desejado + if tipo is None and not _status_in(status_carro, [StatusCarroMapa.CaminhandoRua]): return desejado anterior = np.radians(_safe_float( @@ -631,21 +1049,51 @@ class ControladorMPC: )) dt = _safe_float(dt_controle, 0.5, min_value=0.05, max_value=1.5) - tau = float(self._filtro_angulo_reta_tau_s) - alpha = 1.0 if tau <= 1e-6 else dt / (tau + dt) - filtrado = anterior + alpha * (desejado - anterior) - - max_delta = np.radians(self._taxa_angulo_reta_graus_s) * dt - filtrado = anterior + float(np.clip( - filtrado - anterior, - -max_delta, - max_delta, - )) - dead = np.radians(self._deadband_angulo_reta_graus) + + # Guarda de fase: + # + # Um filtro/rate-limit não pode continuar esterçando para a direita + # quando o MPC já pediu correção significativa para a esquerda, ou + # vice-versa. Nos logs reais isso chegou a persistir por vários ciclos + # e o rover atravessou a centerline antes de o comando cruzar zero. + # + # Na inversão significativa fazemos uma passagem deliberada por zero. + # É mais suave mecanicamente do que saltar direto de +A para -B e, + # principalmente, jamais aplica correção no sentido sabidamente errado. + inversao_significativa = bool( + abs(desejado) > dead + and abs(anterior) > dead + and (desejado * anterior) < 0.0 + ) + + if inversao_significativa: + filtrado = 0.0 + else: + tau = float(self._filtro_angulo_reta_tau_s) + alpha = 1.0 if tau <= 1e-6 else dt / (tau + dt) + filtrado = anterior + alpha * (desejado - anterior) + + max_delta = np.radians(self._taxa_angulo_reta_graus_s) * dt + filtrado = anterior + float(np.clip( + filtrado - anterior, + -max_delta, + max_delta, + )) + if abs(desejado) <= dead and abs(anterior) <= 2.0 * dead: filtrado = 0.0 + # Última barreira: mesmo diante de qualquer alteração futura no filtro, + # uma saída não nula nunca pode permanecer no hemisfério oposto ao + # comando desejado. + if ( + abs(desejado) > dead + and abs(filtrado) > dead + and (desejado * filtrado) < 0.0 + ): + filtrado = 0.0 + return float(np.clip( filtrado, -np.radians(self.angulo_max_graus), @@ -980,23 +1428,37 @@ class ControladorMPC: def _kappa_from_angle(self, tipo, a_rad, L, k_r=0.5, crab_in_phase=False, beta_crab=0.0): """ - Mapeia ângulo de direção 'a_rad' -> curvatura κ (1/m) de acordo com o tipo. - - crab_in_phase=True: movimento 'crab' (quase κ=0), usa desloc. lateral beta_crab (rad). + Curvatura GEOMÉTRICA usada pela LUT do costmap. + + Importante: não usa calcular_omega(v=1). O GPSHandler aplica um fator + dinâmico dependente da velocidade (Ku), então omega(1)/1 não representa + uma curvatura geométrica independente da velocidade. A LUT é deliberadamente + geométrica/conservadora; a simulação pesada do MPC usa calcular_omega com + a velocidade real do ciclo. """ if crab_in_phase: - return 0.0 # tratamos crab fora (x ≈ y * tan(beta_crab)) - if tipo == TipoMovimentoDirecional.RodasDianteiras: - return math.tan(a_rad) / L - if tipo == TipoMovimentoDirecional.RodasTraseiras: - return math.tan(a_rad) / L - elif tipo == TipoMovimentoDirecional.MovimentoArco: - return math.tan((1.0 - k_r) * a_rad) / L - elif tipo == TipoMovimentoDirecional.MovimentoDiagonal: return 0.0 - elif tipo == TipoMovimentoDirecional.MovimentoLateral: + + tipo_enum = _movimento_from_value(tipo) + L_eff = max(float(L), 1e-6) + t = math.tan(float(a_rad)) + + if abs(t) < 1e-9: return 0.0 - else: - return math.tan(a_rad) / L # default + if tipo_enum == TipoMovimentoDirecional.RodasDianteiras: + return t / L_eff + if tipo_enum == TipoMovimentoDirecional.RodasTraseiras: + # Mantém a convenção atualmente usada pelo GPSHandler/C#. + return t / L_eff + if tipo_enum == TipoMovimentoDirecional.MovimentoArco: + # df=+delta, dr=-delta => (tan(df)-tan(dr))/L = 2*tan(delta)/L + return (2.0 * t) / L_eff + if tipo_enum in [ + TipoMovimentoDirecional.MovimentoDiagonal, + TipoMovimentoDirecional.MovimentoLateral, + ]: + return 0.0 + return t / L_eff def _calcula_erro_posicao(self, x, y, ponto_alvo, d_sat=3.0): try: @@ -1007,16 +1469,20 @@ class ControladorMPC: _log(f"Erro ao calcular erro de posicao: {e}") return 1.0, float("inf") - def _calcula_erro_orientacao(self, x, y, theta, ponto_alvo, angulo_caminho, peso_proximo_ponto=0.5): + def _calcula_erro_orientacao(self, x, y, theta, ponto_alvo, angulo_caminho, peso_proximo_ponto=0.5, usar_apenas_caminho=False): try: - orient_sim = self.gps_handler.calcular_orientacao((x, y), ponto_alvo) - orient_sim = (orient_sim + np.pi) % (2 * np.pi) - erro_ori_sim = self.gps_handler.erro_angular(orient_sim, theta, angulo_caminho, peso_proximo_ponto) - #erro_ori_sim = abs(self._wrap_pi(orient_sim - theta_sim)) - #erro_ori_sim = abs(self._wrap_pi(theta_path - theta_sim)) - o_norm = (erro_ori_sim / np.pi) - erro_ori_sim = np.degrees(erro_ori_sim) - return (o_norm, erro_ori_sim) + if usar_apenas_caminho: + erro = abs(self._wrap_pi(float(angulo_caminho) - float(theta))) + else: + # Compatibilidade para manobras: mantém o blend histórico + # entre bearing do alvo e direção local do caminho. + orient_sim = self.gps_handler.calcular_orientacao((x, y), ponto_alvo) + orient_sim = (orient_sim + np.pi) % (2 * np.pi) + erro = self.gps_handler.erro_angular( + orient_sim, theta, angulo_caminho, peso_proximo_ponto + ) + o_norm = abs(float(erro)) / np.pi + return (o_norm, abs(float(np.degrees(erro)))) except Exception as e: _log(f"Erro ao calcular erro de orientacao: {e}") return 1.0, 180.0 @@ -1172,7 +1638,7 @@ class ControladorMPC: (theta_trecho - theta_robo + math.pi) % (2.0 * math.pi) - math.pi ) - erro_maximo = math.radians(35.0 if estrutural else 45.0) + erro_maximo = math.radians(40.0 if estrutural else 50.0) # 35.0 e 45.0 ponto_ficou_para_tras = ( math.isfinite(avanco) @@ -1191,7 +1657,14 @@ class ControladorMPC: return self._proximo_nao_visitado(pontos_visitados) - def _calcular_pesos_movimento(self, contexto, erro_ori_graus, erro_lat_m): + def _calcular_pesos_movimento( + self, + contexto, + erro_ori_graus, + erro_lat_m, + tracking_reta=False, + erro_heading_caminho_graus=None, + ): # helpers simples def _clamp(x, lo, hi): return lo if x < lo else hi if x > hi else x @@ -1241,7 +1714,11 @@ class ControladorMPC: peso_suavidade = _w_from_err_inv(erro_lat_m, w_min=pesos.get("suavidade_min", 0.5), w_max=pesos.get("suavidade_max", 3.6), deadband=0.05, tol=0.25, power=1.3, smooth=True) peso_fator_re = pesos.get("fator_re", 0.05) peso_ideal = pesos.get("ideal", 1.4) - peso_lateral = _w_from_err(erro_lat_m, w_min=pesos.get("lateral_min", 1.8), w_max=pesos.get("lateral_max", 8.5), deadband=0.05, tol=0.7, power=1.3, smooth=True) if dentro else 0.0 + peso_lateral = _w_from_err(erro_lat_m, w_min=pesos.get("lateral_min", 1.8), w_max=pesos.get("lateral_max", 8.5), deadband=self._lateral_deadband_m, tol=0.7, power=1.3, smooth=True) + if not dentro: + # Se saiu da faixa, recuperar a centerline fica MAIS importante, + # nunca menos. A versão anterior zerava este custo. + peso_lateral *= 1.20 # Penalidade proporcional à velocidade penalidade_por_velocidade = (velocidade / self.velocidade_max) / 10.0 @@ -1255,14 +1732,67 @@ class ControladorMPC: erro_ori_abs = abs(float(erro_ori_graus)) erro_lat_abs = abs(float(erro_lat_m)) - if status in [ StatusCarroMapa.EntrandoRua, StatusCarroMapa.SaindoRua, StatusCarroMapa.Manobrando ]: + if tracking_reta: + # V5: os pesos acompanham a mesma política geométrica usada na + # geração de candidatos. Embora normalmente apenas um tipo seja + # gerado por ciclo, manter o custo coerente evita preferências + # contraditórias nas simulações do horizonte. + if erro_heading_caminho_graus is None: + erro_heading_modo = erro_ori_abs + else: + erro_heading_modo = abs(float(erro_heading_caminho_graus)) + + if erro_heading_modo >= self._tracking_arco_entrada_graus: + custo_movimento[TipoMovimentoDirecional.MovimentoArco] += 0.0 + custo_movimento[TipoMovimentoDirecional.RodasDianteiras] += 1.5 + custo_movimento[TipoMovimentoDirecional.MovimentoDiagonal] += 3.0 + elif erro_heading_modo <= self._diagonal_saida_heading_graus: + custo_movimento[TipoMovimentoDirecional.MovimentoArco] += 3.0 + custo_movimento[TipoMovimentoDirecional.MovimentoDiagonal] += 0.0 + custo_movimento[TipoMovimentoDirecional.RodasDianteiras] += 0.25 + else: + custo_movimento[TipoMovimentoDirecional.MovimentoArco] += 1.5 + custo_movimento[TipoMovimentoDirecional.RodasDianteiras] += 0.0 + custo_movimento[TipoMovimentoDirecional.MovimentoDiagonal] += 3.0 + + elif status == StatusCarroMapa.Manobrando: + # Cabeceira/ligação estrutural real: preserva a autoridade do + # Arco. Na geração de candidatos V4 ele continua sendo o único + # tipo permitido neste estado quando não estamos em recuperação + # Rua -> Rua da V3. custo_movimento[TipoMovimentoDirecional.MovimentoArco] += 0.0 - custo_movimento[TipoMovimentoDirecional.RodasDianteiras] += 1.0 - custo_movimento[TipoMovimentoDirecional.MovimentoDiagonal] += 2.0 + custo_movimento[TipoMovimentoDirecional.RodasDianteiras] += 1.5 + custo_movimento[TipoMovimentoDirecional.MovimentoDiagonal] += 3.0 + + elif status in [StatusCarroMapa.EntrandoRua, StatusCarroMapa.SaindoRua]: + # V4: usa o heading do CAMINHO para decidir qual geometria + # merece preferência. O bearing/erro combinado continua nos + # demais custos da manobra, mas não decide sozinho o modo. + if erro_heading_caminho_graus is None: + erro_heading_modo = erro_ori_abs + else: + erro_heading_modo = abs(float(erro_heading_caminho_graus)) + + e_lo = float(self._transicao_rua_dianteira_heading_graus) + e_hi = float(self._transicao_rua_arco_heading_graus) + + t = (erro_heading_modo - e_lo) / max(e_hi - e_lo, 1e-6) + t = _smoothstep01(_clamp(t, 0.0, 1.0)) + + penal = float(self._transicao_rua_penalidade_modo) + + # t=0: alinhado -> dianteira preferida + # t=1: desalinhado -> arco preferido + custo_movimento[TipoMovimentoDirecional.RodasDianteiras] += penal * t + custo_movimento[TipoMovimentoDirecional.MovimentoArco] += penal * (1.0 - t) + + # Diagonal não pertence a esta política de transição. + custo_movimento[TipoMovimentoDirecional.MovimentoDiagonal] += 3.0 + else: - custo_movimento[TipoMovimentoDirecional.MovimentoArco] += 0.0 if erro_ori_abs >= 10.0 else 1.0 - custo_movimento[TipoMovimentoDirecional.RodasDianteiras] += 1.2 if erro_ori_abs >= 10.0 else 0.0 - custo_movimento[TipoMovimentoDirecional.MovimentoDiagonal] += 0.0 if erro_ori_abs <= 2.0 and erro_lat_abs <= 0.4 else 2.0 + custo_movimento[TipoMovimentoDirecional.MovimentoArco] += 1.5 + custo_movimento[TipoMovimentoDirecional.RodasDianteiras] += 0.0 + custo_movimento[TipoMovimentoDirecional.MovimentoDiagonal] += 0.5 return (peso_erro_pos, peso_erro_ori, peso_suavidade, peso_fator_re, peso_ideal, peso_lateral, custo_movimento) except Exception as e: @@ -1769,7 +2299,11 @@ class ControladorMPC: #candidatos_ativos = self._selecionar_melhores(candidatos_ativos, N=N_top) beam_topN_max = int(getattr(self, "_beam_topN_max", 5)) beam_topN_min = int(getattr(self, "_beam_topN_min", 2)) - candidatos_ativos = self._filtrar_por_margem_angular(novos_candidatos, margem_graus=2.0) + candidatos_ativos = self._filtrar_por_margem_angular( + novos_candidatos, + margem_graus=2.0, + por_tipo=True, + ) # número alvo de trajetórias ativas #N_alvo = min(beam_topN_max, max(beam_topN_min, beam_traj_min)) beam_traj_min = max(1, min(beam_traj_min, beam_topN_max)) # sanity @@ -1814,15 +2348,24 @@ class ControladorMPC: status_carro, comando_anterior, self.tempo_execucao_local, + tipo_movimento=tipo_final, ) ponto_alvo_primario = ponto_alvo_real["xy"] - erro_lateral, _, _, _, theta_path, _ = self.cross_track_error_point( - x, - y, - idx_referencia=idx_alvo_correcao, + ref_saida = self._referencia_caminho_lookahead( + x, y, idx_alvo_correcao, velocidade ) - _, erro_orientacao = self._calcula_erro_orientacao(x, y, theta, ponto_alvo_primario, theta_path, 0.5) + erro_lateral = float(ref_saida["e_lat"]) + erro_heading_saida = self._erro_heading_caminho(theta, ref_saida["theta_ref"]) + tracking_reta_saida = self._tracking_reta_permitido( + contexto, idx_alvo_correcao, erro_lateral, erro_heading_saida + ) + if tracking_reta_saida: + erro_orientacao = abs(float(np.degrees(erro_heading_saida))) + else: + _, erro_orientacao = self._calcula_erro_orientacao( + x, y, theta, ponto_alvo_primario, ref_saida["theta_local"], 0.5 + ) _cmd = { "enviar_comando": True, @@ -1837,6 +2380,11 @@ class ControladorMPC: "idx_alvo_autoritativo": int(idx_alvo_correcao), "alvo_direcional_xy": ponto_alvo_primario, "lookahead_direcional_m": float(ponto_alvo_real.get("_lookahead_m", 0.0)), + "tracking_reta": bool(tracking_reta_saida), + "heading_caminho_lookahead_graus": float(np.degrees(ref_saida["theta_ref"])), + "heading_caminho_local_graus": float(np.degrees(ref_saida["theta_local"])), + "heading_lookahead_m": float(ref_saida["lookahead_heading_m"]), + "erro_heading_caminho_graus": abs(float(np.degrees(erro_heading_saida))), "debug_custo": debug_custo, "candidatos_testados": K, "motivos": [] @@ -1946,61 +2494,129 @@ class ControladorMPC: """ try: now = time.perf_counter + CARRO = _as_dict(contexto.get("Carro", {})) + status = CARRO.get("Status", StatusCarroMapa.Parado.value) + velocidade = max(0.0, _safe_float(CARRO.get("Velocidade", 0.0), 0.0)) + idx_base = _safe_int(ponto_alvo.get("_idx_base", 0), 0) - # 1) Ângulo ideal (wrap [-pi,pi] e clamp ao máx) - orient = self.gps_handler.calcular_orientacao(ponto_alvo["xy"], ponto_atual) - delta_theta = (orient - ponto_atual[2] + np.pi) % (2*np.pi) - np.pi - - #delta_theta, theta_ref, dbg = self.delta_por_dist_esq_dir( - # estado_atual=ponto_atual, - # ponto_alvo_xy=ponto_alvo["xy"], - # orient_corredor=np.radians(contexto["Carro"]["AnguloCaminho"]), # em rad - # dist_esq=contexto["Carro"]["DistanciaEsquerda"], - # dist_dir=contexto["Carro"]["DistanciaDireita"], - # angulo_max_graus=self.angulo_max_graus, - # sigma_factor=0.6, - # w_min=0.15, - # usar_nudge_centro=True, - # k_e=1.0, v=contexto["Carro"].get("Velocidade", 0.0), v0=0.3 - #) - - ang_max_rad = np.radians(self.angulo_max_graus) - angulo_ideal_rad = float(np.clip(delta_theta, -ang_max_rad, ang_max_rad)) - angulo_ideal_deg = float(np.degrees(angulo_ideal_rad)) - #print(f"Orient: {np.degrees(orient)}, delta theta: {np.degrees(delta_theta)}, angulo ideal: {angulo_ideal_deg}") - - e_lat, _, _, _, theta_path, _ = self.cross_track_error_point( - ponto_atual[0], - ponto_atual[1], - idx_referencia=ponto_alvo.get("_idx_base"), + # -------------------------------------------------------------- + # PATH TRACKING: heading do CAMINHO + cross-track, sem mirar ponto. + # A geometria é escolhida por erro de heading/cross-track (V5). + # MANOBRA estrutural: mantém pure-pursuit e Arco obrigatório. + # -------------------------------------------------------------- + ref_path = self._referencia_caminho_lookahead( + ponto_atual[0], ponto_atual[1], idx_base, velocidade ) - (_, e_ori) = self._calcula_erro_orientacao(ponto_atual[0], ponto_atual[1], ponto_atual[2], ponto_alvo["xy"], theta_path, 0.5) + e_lat = float(ref_path["e_lat"]) + theta_ref = float(ref_path["theta_ref"]) + erro_heading = self._erro_heading_caminho(ponto_atual[2], theta_ref) + e_ori = abs(math.degrees(erro_heading)) - # 2) Tipos válidos conforme contexto - CARRO = contexto.get("Carro", {}) - status = CARRO.get("Status", 0) - if _status_in(status, [StatusCarroMapa.EntrandoRua, StatusCarroMapa.SaindoRua, StatusCarroMapa.Manobrando]): - tipos_validos = [TipoMovimentoDirecional.MovimentoArco] - elif abs(e_ori) > 15.0: - tipos_validos = [TipoMovimentoDirecional.RodasDianteiras, TipoMovimentoDirecional.MovimentoArco] - else: - # Histerese de modo na passada. Antes, 0,20 m/3 graus era - # uma fronteira seca: qualquer ruido alternava Diagonal e - # RodasDianteiras. Entrar exige uma faixa estreita; depois de - # entrar, a faixa de permanencia e um pouco maior. - tipo_preferido_enum = _movimento_from_value(tipo_preferido) - manter_diagonal = ( - tipo_preferido_enum == TipoMovimentoDirecional.MovimentoDiagonal - and abs(e_ori) <= 4.0 - and abs(e_lat) <= 0.30 + tracking_reta = self._tracking_reta_permitido( + contexto, idx_base, e_lat, erro_heading + ) + + if tracking_reta: + tipo_ref, angulo_ideal_rad = self._referencia_direcional_reta( + e_lat, + erro_heading, + velocidade, + tipo_preferido=tipo_preferido, ) - entrar_diagonal = abs(e_ori) <= 1.5 and abs(e_lat) <= 0.12 + tipos_validos = [tipo_ref] + else: + orient = self.gps_handler.calcular_orientacao(ponto_alvo["xy"], ponto_atual) + delta_theta = self._wrap_pi(orient - ponto_atual[2]) + angulo_ideal_rad = float(np.clip( + delta_theta, + -np.radians(self.angulo_max_graus), + np.radians(self.angulo_max_graus), + )) + + tipo_preferido_enum = _movimento_from_value(tipo_preferido) + erro_heading_alvo_deg = abs(float(math.degrees(delta_theta))) + + if _status_in(status, [StatusCarroMapa.Manobrando]): + # Manobra estrutural real continua com Arco obrigatório. + # A recuperação Rua -> Rua da V3 entra antes em + # tracking_reta=True, portanto não cai neste ramo. + tipos_validos = [TipoMovimentoDirecional.MovimentoArco] + + elif _status_in( + status, + [StatusCarroMapa.EntrandoRua, StatusCarroMapa.SaindoRua], + ): + # V4: entrada/saída são estados de TRANSIÇÃO. Arco e + # RodasDianteiras são avaliados pelo mesmo MPC. + # + # Colocar o tipo anterior primeiro só estabiliza a divisão + # do orçamento quando sobra um candidato ímpar. A escolha + # final continua sendo pelo custo acumulado. + if tipo_preferido_enum == TipoMovimentoDirecional.RodasDianteiras: + tipos_validos = [ + TipoMovimentoDirecional.RodasDianteiras, + TipoMovimentoDirecional.MovimentoArco, + ] + else: + tipos_validos = [ + TipoMovimentoDirecional.MovimentoArco, + TipoMovimentoDirecional.RodasDianteiras, + ] + + elif _status_in(status, [StatusCarroMapa.Direcionando]): + # Fora da rua, um heading muito errado deve ser corrigido + # com Arco, que possui raio de giro menor. Quando o rover + # já estiver alinhado, RodasDianteiras finaliza de forma + # mais suave. A faixa 8°..15° é a histerese: se já entrou + # em Arco, permanece nele até cruzar o limiar de saída. + manter_arco = ( + tipo_preferido_enum == TipoMovimentoDirecional.MovimentoArco + and erro_heading_alvo_deg > self._direcionando_arco_saida_graus + ) + + if ( + erro_heading_alvo_deg >= self._direcionando_arco_entrada_graus + or manter_arco + ): + tipos_validos = [TipoMovimentoDirecional.MovimentoArco] + else: + tipos_validos = [TipoMovimentoDirecional.RodasDianteiras] + + elif _status_in(status, [StatusCarroMapa.RetornandoBase]): + # Guarda defensiva. Em V5 RetornandoBase deve entrar no + # path tracking acima; se algum contexto incompleto impedir + # isso, ainda assim nunca volta ao antigo fallback Front-only. + manter_arco = ( + tipo_preferido_enum == TipoMovimentoDirecional.MovimentoArco + and erro_heading_alvo_deg > self._tracking_arco_saida_graus + ) + if ( + erro_heading_alvo_deg >= self._tracking_arco_entrada_graus + or manter_arco + ): + tipos_validos = [TipoMovimentoDirecional.MovimentoArco] + else: + tipos_validos = [TipoMovimentoDirecional.RodasDianteiras] - if manter_diagonal or entrar_diagonal: - tipos_validos = [TipoMovimentoDirecional.MovimentoDiagonal] else: tipos_validos = [TipoMovimentoDirecional.RodasDianteiras] + # Para ranking/telemetria de manobra, preserva erro histórico. + _, _, _, _, theta_local, _ = self.cross_track_error_point( + ponto_atual[0], ponto_atual[1], idx_referencia=idx_base + ) + _, e_ori = self._calcula_erro_orientacao( + ponto_atual[0], ponto_atual[1], ponto_atual[2], + ponto_alvo["xy"], theta_local, 0.5, usar_apenas_caminho=False + ) + + angulo_ideal_rad = float(np.clip( + angulo_ideal_rad, + -np.radians(self.angulo_max_graus), + np.radians(self.angulo_max_graus), + )) + angulo_ideal_deg = float(np.degrees(angulo_ideal_rad)) + # 3) Flags de matriz if dados_costmap is None: dados_costmap = self._get_costmap_direcional(contexto) @@ -2008,6 +2624,11 @@ class ControladorMPC: # 4) Geração bruta (graus) angulos_raw = self._gerar_candidatos_brutos(angulo_ideal_deg, filtrar_matriz) + if tipos_validos == [TipoMovimentoDirecional.MovimentoDiagonal]: + lim_diag = float(min(self._angulo_diagonal_max_graus, self.angulo_max_graus)) + angulos_raw = [a for a in angulos_raw if abs(float(a)) <= lim_diag + 1e-9] + if not angulos_raw: + angulos_raw = [float(np.clip(angulo_ideal_deg, -lim_diag, lim_diag))] if not angulos_raw: # fallback mínimo em torno do ideal base = [angulo_ideal_deg + d for d in (-2.0, -1.0, 0.0, 1.0, 2.0)] @@ -2449,21 +3070,25 @@ class ControladorMPC: omega = float(self.gps_handler.calcular_omega(v, angulo, tipo)) k = np.arange(1, n + 1, dtype=np.float32) - if abs(omega) < 1e-6: - # reta: mesmo esquema do seu _nova_posicao - dx = dist_step * k * math.sin(th0) - dy = dist_step * k * math.cos(th0) + if tipo in [TipoMovimentoDirecional.MovimentoDiagonal, TipoMovimentoDirecional.MovimentoLateral]: + # 4WS em fase: heading do corpo permanece, mas a velocidade + # translacional aponta para theta + angulo. + direcao = th0 + float(angulo) + dx = dist_step * k * math.sin(direcao) + dy = dist_step * k * math.cos(direcao) th = th0 + np.zeros_like(k, dtype=np.float32) x = x0 + dx y = y0 + dy else: - R = v / omega + # Mesma integração semi-implícita usada por _nova_posicao e + # pelo simulador C#: primeiro atualiza theta, depois translada. + # O cumsum reproduz exatamente a aplicação passo a passo. dth = omega * dt - th = th0 + dth * k - s0, c0 = math.sin(th0), math.cos(th0) - # solução fechada consistente com x+=sin(θ)*d, y+=cos(θ)*d - x = x0 + R * (c0 - np.cos(th)) - y = y0 + R * (np.sin(th) - s0) + th = th0 + dth * k + dx_step = dist_step * np.sin(th) + dy_step = dist_step * np.cos(th) + x = x0 + np.cumsum(dx_step) + y = y0 + np.cumsum(dy_step) return list(zip(x.astype(float).tolist(), y.astype(float).tolist(), @@ -2841,9 +3466,16 @@ class ControladorMPC: sub_dt = dt_total / n_subs dist_sub = v_planejado * sub_dt # custo por metro - # --- PATCH: estado local p/ amortecimento e anti-chatter - e_lat_prev_local = None - side_mem_local = 0 # -1 direita, +1 esquerda, 0 centro + # Estado inicial real do candidato para derivada/anti-chatter. + try: + e_lat_inicial, _, _, _, _, _ = self.cross_track_error_point( + x_sim, y_sim, idx_referencia=self._proximo_nao_visitado(pontos_visitados) + ) + e_lat_prev_local = float(e_lat_inicial) + side_mem_local = 0 if abs(e_lat_inicial) < max(0.08, self._lateral_deadband_m * 2.0) else (1 if e_lat_inicial > 0 else -1) + except Exception: + e_lat_prev_local = None + side_mem_local = 0 for n_sub in range(n_subs): # IMPORTANTE: _nova_posicao precisa aceitar dt opcional @@ -2875,90 +3507,109 @@ class ControladorMPC: # -------------------------- # ERROS / CUSTOS # -------------------------- - # posição (mantido) - (d_norm, erro_pos) = self._calcula_erro_posicao(x_sim, y_sim, ponto_alvo_sim) - - # --- PATCH A: usa theta do CAMINHO (segmento projetado) como referência - # e usa também margem local pra normalização do custo lateral - e_lat, _, _, _, theta_path, margem = self.cross_track_error_point( - x_sim, - y_sim, - idx_referencia=idx_alvo_sim, + ref_path_sim = self._referencia_caminho_lookahead( + x_sim, y_sim, idx_alvo_sim, v_planejado + ) + e_lat = float(ref_path_sim["e_lat"]) + theta_path = float(ref_path_sim["theta_ref"]) + margem = float(ref_path_sim["margem"]) + erro_heading_sim = self._erro_heading_caminho(theta_sim, theta_path) + tracking_reta_sim = self._tracking_reta_permitido( + contexto, idx_alvo_sim, e_lat, erro_heading_sim ) - # erro de orientação - (o_norm, erro_ori_sim) = self._calcula_erro_orientacao(x_sim, y_sim, theta_sim, ponto_alvo_sim, theta_path, 0.5) + if tracking_reta_sim: + # Passada: NÃO existe atração para ponto. O custo é path + # tracking puro: heading da centerline + cross-track. + erro_pos = abs(e_lat) + d_norm = 0.0 + erro_ori_sim = abs(float(np.degrees(erro_heading_sim))) + o_norm = abs(float(erro_heading_sim)) / np.pi + else: + (d_norm, erro_pos) = self._calcula_erro_posicao( + x_sim, y_sim, ponto_alvo_sim + ) + (o_norm, erro_ori_sim) = self._calcula_erro_orientacao( + x_sim, y_sim, theta_sim, + ponto_alvo_sim, ref_path_sim["theta_local"], 0.5, + usar_apenas_caminho=False, + ) - # pesos dinâmicos (mantido) - pesos = self._calcular_pesos_movimento(contexto, erro_ori_sim, e_lat) - (peso_erro_pos, peso_erro_ori, peso_suavidade, peso_fator_re, peso_ideal, peso_lateral, peso_movimento) = pesos + pesos = self._calcular_pesos_movimento( + contexto, + erro_ori_sim, + e_lat, + tracking_reta=tracking_reta_sim, + erro_heading_caminho_graus=abs( + float(np.degrees(erro_heading_sim)) + ), + ) + (peso_erro_pos, peso_erro_ori, peso_suavidade, + peso_fator_re, peso_ideal, peso_lateral, + peso_movimento) = pesos - # ângulo relativo ao alvo e suavidade (parcialmente mantido) - vetor_alvo = np.array(ponto_alvo_sim) - np.array([x_sim, y_sim]) - vetor_alvo_norm = vetor_alvo / (np.linalg.norm(vetor_alvo) + 1e-9) - # Mesmo referencial de _nova_posicao: x=Leste usa sin(theta) - # e y=Norte usa cos(theta). - vetor_movel = np.array([np.sin(theta_sim), np.cos(theta_sim)]) - cos_delta = float(np.dot(vetor_movel, vetor_alvo_norm)) - delta_angulo = abs(angulo_testado - prev_ang) + delta_angulo = abs(self._wrap_pi(float(angulo_testado) - float(prev_ang))) # -------------------------- # COMPONENTES DE CUSTO # -------------------------- - # posição (mantido) - #custo_pos = erro_pos * peso_erro_pos - custo_pos = d_norm * peso_erro_pos - - # --- PATCH B: custo de orientação desacoplado de erro_pos e quadrático - #custo_ori = erro_ori_sim * peso_erro_ori + custo_pos = 0.0 if tracking_reta_sim else d_norm * peso_erro_pos custo_ori = o_norm * peso_erro_ori - # --- PATCH C: suavidade SEM depender de erro_pos + piso - base_suav = 0.25 # piso de penalização (0.25–0.5) - erro_suavidade = ((base_suav + 1.0) * (abs(delta_angulo) / np.pi) * erro_pos) + # Suavidade realmente independente do erro de posição. + ang_norm = delta_angulo / max(np.radians(self.angulo_max_graus), 1e-6) + erro_suavidade = float(ang_norm * ang_norm) custo_suavidade = erro_suavidade * peso_suavidade - # tipo de movimento (mantido) custo_tipo_movimento = float(peso_movimento[tipo]) + if prev_tipo != tipo: + custo_tipo_movimento += float(self._penalidade_troca_tipo) - # fator "re" (mantido como estava) - erro_re = -cos_delta # já está em [-1, +1] - custo_re = erro_re * peso_fator_re # aplica peso normalmente + if tracking_reta_sim: + # Pequena recompensa por manter o corpo paralelo à linha; + # não há vetor carro->ponto neste termo. + cos_delta = math.cos(float(erro_heading_sim)) + else: + vetor_alvo = np.array(ponto_alvo_sim) - np.array([x_sim, y_sim]) + vetor_alvo_norm = vetor_alvo / (np.linalg.norm(vetor_alvo) + 1e-9) + vetor_movel = np.array([np.sin(theta_sim), np.cos(theta_sim)]) + cos_delta = float(np.dot(vetor_movel, vetor_alvo_norm)) + erro_re = -float(cos_delta) + custo_re = erro_re * peso_fator_re - # --- PATCH D: custo lateral com deadband + quadrático + normalização - dead = 0.0 # 5–12 cm - half = max(float(margem) / 2.0, 0.5) # meia-largura mínima + # Cross-track sempre ativo, inclusive fora do corredor. + dead = self._lateral_deadband_m + half = max(float(margem) / 2.0, 0.5) e_eff = max(0.0, abs(e_lat) - dead) - r = e_eff/half + r = e_eff / half if r <= 1.0: - alpha = 0.25 # 0 = só quadrático; 1 = só linear + alpha = 0.25 l_norm = alpha*r + (1-alpha)*r**2 else: - gamma = 1.0 - l_norm = 1.0 + gamma*(r - 1.0) # ex.: gamma = 1.0 (ajuste) + l_norm = 1.0 + (r - 1.0) custo_lateral = l_norm * peso_lateral - # --- PATCH E: amortecimento por derivada do erro lateral (m/s) + # Derivada funciona mesmo com n_subporpasso=1 porque o estado + # inicial do candidato é usado como memória. if e_lat_prev_local is not None: - de = (e_lat - e_lat_prev_local) / max(sub_dt, 1e-6) # m/s - peso_lat_der = 0.1 * float(peso_lateral) # ganho pequeno + de = (e_lat - e_lat_prev_local) / max(sub_dt, 1e-6) + peso_lat_der = 0.08 * float(peso_lateral) custo_lat_der = (de ** 2) * peso_lat_der - - dlat = (abs(e_lat) - abs(e_lat_prev_local)) / max(sub_dt,1e-6) - pen_away = max(0.0, dlat/half) * (0.2 * peso_lateral) + dlat = (abs(e_lat) - abs(e_lat_prev_local)) / max(sub_dt, 1e-6) + pen_away = max(0.0, dlat / half) * (0.15 * peso_lateral) else: custo_lat_der = 0.0 pen_away = 0.0 custo_lateral += pen_away e_lat_prev_local = e_lat - # --- PATCH F: anti-chatter (histerese) perto do centro - band = 0.10 # 8–12 cm + band = max(0.08, self._lateral_deadband_m * 2.0) side_now = 0 if abs(e_lat) < band else (1 if e_lat > 0 else -1) flip_pen = 0.0 - if abs(e_lat) < band and side_mem_local != 0 and side_now != side_mem_local: - flip_pen = 0.12 * float(peso_suavidade) # 0.08–0.20 - side_mem_local = side_now + if side_mem_local != 0 and side_now != 0 and side_now != side_mem_local: + flip_pen = 0.12 * float(peso_suavidade) + if side_now != 0: + side_mem_local = side_now # soma local (heurística do visual + demais) e NORMALIZA por metro custo_mapa = ( diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/camera_manager.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/camera_manager.py index cb0dea483..ce16e08d8 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/camera_manager.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/camera_manager.py @@ -4,16 +4,16 @@ import datetime import os import threading import time -from typing import Any, Dict, Optional +from typing import Optional, Set import cv2 import numpy as np from camera_worker.camera_oak import CameraOak -from camera_worker.segformer_runner import SegformerNavRunner +from visual_worker.inferencia.corridor_onnx_runner import CorridorOnnxRunner from shared.contexto_global_redis import ContextoGlobalRedis, CtxKey -from shared.enums import StatusModulo, T_Code, TipoFrameCamera +from shared.enums import StatusModulo, StatusCarroMapa, T_Code, TipoFrameCamera from shared.gpu_priority_controller import GpuPriorityController from shared.perf_monitor import VisualPerfMonitor from shared.utils import get_velocidade_atual_ms @@ -27,7 +27,6 @@ from visual_worker.processamento.costmap_fuser import CostmapFuser from visual_worker.processamento.visual_grid_builder import ( GridGeometryConfig, GridReferenceBuilder, - GridConfidenceConfig, VisualGridBuilder, ) from visual_worker.processamento.visual_debug_renderer import VisualDebugRenderer @@ -35,12 +34,12 @@ from visual_worker.processamento.visual_debug_renderer import VisualDebugRendere class CameraManager: """ - CameraManager v1 do Visual Worker. + CameraManager v2 do Visual Worker. Responsabilidade: - Inicializar OAK-D Lite. - - Inicializar SegformerNavRunner ONNX/TensorRT. - - Inicializar SegmentacaoManager v1. + - Inicializar CorridorOnnxRunner ONNX/TensorRT. + - Inicializar SegmentacaoManager v2. - Inicializar grid builder + CostmapFuser. - Orquestrar loops: segmentação -> detecção onboard -> grid/costmap -> publicação -> stream/debug. @@ -83,7 +82,12 @@ class CameraManager: self.operante = False self.iniciando = False self.debug_visual = False - self.debug_perf = False + + # Telemetria técnica fica explicitamente controlada pelo config. + self.telemetry_enabled = True + self.telemetry_console = True + self.telemetry_publish_runtime = True + self.telemetry_runner_timing = True self.largura_robo_m = 0.85 @@ -92,10 +96,13 @@ class CameraManager: self._loop_grid_iniciado = False self._loop_publicacao_iniciado = False self._loop_stream_iniciado = False + self._loop_preview_iniciado = False self._loop_analise_iniciado = False self._vida_lock = threading.RLock() self._cache_lock = threading.RLock() + self._preview_lock = threading.RLock() + self._preview_cond = threading.Condition(self._preview_lock) self._pub_lock = threading.RLock() self._fechando_camera = False @@ -110,6 +117,13 @@ class CameraManager: self._posproc_salvando = False self.posproc_intervalo_min_s = 5.0 + # Preview é um fluxo lateral ao runtime crítico. Ele só é produzido + # quando algum consumidor solicita um tipo de frame. + self.preview_fps = 2.0 + self.preview_request_timeout_s = 0.35 + self.preview_cache_fresh_s = 0.75 + self.replay_jpeg_quality = 70 + self._reset_runtime_state() # ============================================================ @@ -121,7 +135,7 @@ class CameraManager: self._ultimo_depth_frame = None self._ultimo_predictions = None - self._ultimo_seg_aux_result = None + self._ultimo_status_result = None self._ultima_analise_segmentacao = None self._ultimo_detections = [] @@ -132,24 +146,35 @@ class CameraManager: self._startup_grace_until = 0.0 self._seg_cache = { + # ts mantido como alias de compatibilidade para result_ts. "ts": 0.0, + "result_ts": 0.0, "frame_ts": 0.0, "predictions": None, - "aux_result": None, + "status_result": None, "analise": None, "res": None, "infer_ms": 0.0, "post_ms": 0.0, + "runner_timing_ms": {}, } self._det_cache = { + # ts mantido como alias de compatibilidade para result_ts. "ts": 0.0, + "result_ts": 0.0, + "frame_ts": 0.0, "detections": [], "res": None, } self._grid_cache = { + # ts mantido como alias de compatibilidade para result_ts. "ts": 0.0, + "result_ts": 0.0, + "depth_frame_ts": 0.0, + "seg_frame_ts": 0.0, + "det_frame_ts": 0.0, "snapshot": None, "grid_conf": None, } @@ -182,6 +207,7 @@ class CameraManager: self._ultimo_frame_preview_por_tipo = {} self._ultimo_frame_preview_ts_por_tipo = {} + self._preview_requests: Set[str] = set() def inicializar(self, mx_id): if self.iniciando: @@ -206,9 +232,18 @@ class CameraManager: self.det_config = load_det_config() self.posproc_intervalo_min_s = float(self.seg_config.get("posproc_intervalo_min_s", 5.0)) + self.preview_fps = float(self.seg_config.get("preview_fps", 2.0)) + self.preview_request_timeout_s = float(self.seg_config.get("preview_request_timeout_s", 0.35)) + self.preview_cache_fresh_s = float(self.seg_config.get("preview_cache_fresh_s", 0.75)) + self.replay_jpeg_quality = int(self.seg_config.get("replay_jpeg_quality", 70)) self.debug_visual = bool(self.seg_config.get("debug_visual", False)) - self.debug_perf = bool(self.seg_config.get("debug_perf", False)) + + telemetry_cfg = self.seg_config.get("telemetry", {}) or {} + self.telemetry_enabled = bool(telemetry_cfg.get("enabled", True)) + self.telemetry_console = bool(telemetry_cfg.get("console_performance", True)) + self.telemetry_publish_runtime = bool(telemetry_cfg.get("publish_runtime_info", True)) + self.telemetry_runner_timing = bool(telemetry_cfg.get("runner_detailed_timing", True)) self._inicializar_camera(mx_id) @@ -278,17 +313,37 @@ class CameraManager: self.mostrar_log(f"⚠️ Camera com ID {mx_id} não conectada: {e}") def _inicializar_segmentador(self): - self.seg_runner = SegformerNavRunner(self.seg_config) + model_path = self.seg_config.get("model_path") + if not model_path: + raise RuntimeError("Visual Worker v2 sem model_path do ONNX de corredor.") + self.seg_runner = CorridorOnnxRunner( + onnx_path=model_path, + runtime_config=self.seg_config.get("model_runtime", {}) or {}, + telemetry_config=self.seg_config.get("telemetry", {}) or {}, + mostrar_log=self.mostrar_log, + ) + + # Contrato ONNX é a fonte da verdade para classes e status. + self.seg_runner.contract.validate_status_enum(StatusCarroMapa) + + runtime = self.seg_runner.runtime_info() self.mostrar_log( - "[SEG_ONNX] Runner iniciado | " - f"model={self.seg_config.get('onnx_model_path')} | " - f"provider={self.seg_config.get('onnx_provider', 'tensorrt')} | " - f"resolution={self.seg_config.get('ia_resolution')}" + "[SEG_RUNTIME] pronto | " + f"model={runtime['model_path']} | " + f"provider={runtime['provider_primary']} | " + f"resolution={self.seg_runner.contract.resolution_wh}" ) def _inicializar_analisador_segmentacao(self): - cfg_dict = self.seg_config.get("segmentacao", {}) or {} + cfg_dict = dict(self.seg_config.get("segmentacao", {}) or {}) + + # Nunca duplicar estes IDs no config. Eles pertencem ao modelo. + cfg_dict["id_navegavel"] = int(self.seg_runner.contract.nav_class_id) + cfg_dict["id_nao_navegavel"] = int(self.seg_runner.contract.non_nav_class_id) + cfg_dict["include_model_probs"] = bool( + (self.seg_config.get("telemetry", {}) or {}).get("include_status_probs", False) + ) seg_cfg = SegmentacaoConfig(**{ k: v @@ -298,16 +353,20 @@ class CameraManager: self.segmentacao_manager = SegmentacaoManager( config=seg_cfg, - color_map=getattr(self.seg_runner, "colormap_rgb", None), - classes=getattr(self.seg_runner, "classes", None), + color_map=self.seg_runner.colormap_rgb, + classes=self.seg_runner.classes, ) def _inicializar_grid(self): grid_cfg = self.seg_config.get("grid", {}) or {} geom_cfg = grid_cfg.get("geometry", {}) or {} - conf_cfg = grid_cfg.get("confidence", {}) or {} + conf_cfg = dict(grid_cfg.get("confidence", {}) or {}) fuser_cfg = grid_cfg.get("fuser", {}) or {} + # IDs semânticos pertencem ao contrato do modelo, não ao config. + conf_cfg["id_navegavel"] = int(self.seg_runner.contract.nav_class_id) + conf_cfg["id_nao_navegavel"] = int(self.seg_runner.contract.non_nav_class_id) + if "grid_shape" not in geom_cfg: geom_cfg["grid_shape"] = tuple(grid_cfg.get("grid_shape", (15, 10))) @@ -388,6 +447,10 @@ class CameraManager: freq_inferencia = float(seg_config.get("inferencia_fps", freq_analise)) freq_deteccao = float(seg_config.get("deteccao_fps", freq_analise)) + if not self._loop_preview_iniciado: + self._iniciar_loop_preview(freq=self.preview_fps) + self._loop_preview_iniciado = True + if not self._loop_stream_iniciado: stream = getattr(self.camera, "stream", None) stream_fps = float(getattr(stream, "_op_fps", 2.0) or 2.0) @@ -460,7 +523,7 @@ class CameraManager: depth_frame = None depth_ts = None predictions = None - aux_result = None + status_result = None analise_seg = None dets = [] @@ -511,28 +574,32 @@ class CameraManager: for i in range(max(0, n_seg)): try: rgb_frame, frame_ts, res = self.get_rgb_frame() - rgb_frame = self._validar_frame_rgb_para_segmentacao(rgb_frame) + rgb_frame = self._validar_frame_camera_para_segmentacao(rgb_frame) if rgb_frame is None: continue t_inf0 = time.perf_counter() - predictions, ts, roi_resized, roi_info, aux_result = self.seg_runner.infer_ids(rgb_frame) + infer_result = self.seg_runner.infer(rgb_frame) t_inf1 = time.perf_counter() - if predictions is None or ts is None: - continue + predictions = infer_result.seg_ids + status_result = infer_result.status t_post0 = time.perf_counter() - analise_seg = self.segmentacao_manager.analisar(predictions, aux_result=aux_result) + analise_seg = self.segmentacao_manager.analisar( + predictions, + status_result=status_result, + ) t_post1 = time.perf_counter() self._store_segmentacao( predictions=predictions, - aux_result=aux_result, + status_result=status_result, analise=analise_seg, frame_ts=frame_ts, res=res, + runner_timing_ms=infer_result.timing_ms, infer_ms=(t_inf1 - t_inf0) * 1000.0, post_ms=(t_post1 - t_post0) * 1000.0, ) @@ -561,7 +628,7 @@ class CameraManager: if depth_frame is None: continue - snapshot, grid_conf = self._processar_grid( + snapshot, grid_conf, _grid_timing = self._processar_grid( depth_frame=depth_frame, segmentacao=predictions, dados_visuais_seg=analise_seg.get("dados_visuais", {}), @@ -572,7 +639,13 @@ class CameraManager: if snapshot is None: continue - self._store_grid(snapshot=snapshot, grid_conf=grid_conf) + self._store_grid( + snapshot=snapshot, + grid_conf=grid_conf, + depth_frame_ts=depth_ts or 0.0, + seg_frame_ts=float(self._seg_cache.get("frame_ts", 0.0) or 0.0), + det_frame_ts=float(self._det_cache.get("frame_ts", 0.0) or 0.0), + ) grid_ok += 1 @@ -628,21 +701,25 @@ class CameraManager: return False - def _validar_frame_rgb_para_segmentacao(self, rgb_frame): - if rgb_frame is None: + def _validar_frame_camera_para_segmentacao(self, frame_camera): + if frame_camera is None: return None - if not hasattr(rgb_frame, "shape") or rgb_frame.size == 0: + if not hasattr(frame_camera, "shape") or frame_camera.size == 0: return None - if rgb_frame.ndim != 3 or rgb_frame.shape[2] != 3: - self.mostrar_log(f"[SEG_ONNX] RGB inválido shape={getattr(rgb_frame, 'shape', None)}") + if frame_camera.ndim != 3 or frame_camera.shape[2] != 3: + self.mostrar_log(f"[SEG_ONNX] Frame câmera inválido shape={getattr(frame_camera, 'shape', None)}") return None - if rgb_frame.dtype != np.uint8: - rgb_frame = np.clip(rgb_frame, 0, 255).astype(np.uint8) + if frame_camera.dtype != np.uint8: + self.mostrar_log( + f"[SEG_ONNX] Frame câmera precisa ser uint8; veio dtype={frame_camera.dtype}" + ) + return None - return np.ascontiguousarray(rgb_frame) + # Normalmente o frame da OAK já é C-contiguous. Só copia se a fonte mudar. + return frame_camera if frame_camera.flags.c_contiguous else np.ascontiguousarray(frame_camera) def get_rgb_frame(self): cam = self.camera @@ -651,7 +728,9 @@ class CameraManager: return None, None, None try: - frame, res = cam.requisitar_frame_rgb() + # Hot path do Visual Worker: referência read-only ao latest frame. + # CameraOak substitui o cache, nunca modifica o ndarray antigo. + frame, res = cam.requisitar_frame_rgb_ref() erro = (res or {}).get("erro") if self._is_erro_fatal_camera(erro): @@ -684,7 +763,7 @@ class CameraManager: return None, None, None try: - frame, res = cam.requisitar_frame_depth() + frame, res = cam.requisitar_frame_depth_ref() erro = (res or {}).get("erro") if self._is_erro_fatal_camera(erro): @@ -735,50 +814,83 @@ class CameraManager: # Store helpers # ============================================================ - def _store_segmentacao(self, predictions, aux_result, analise, frame_ts, res, infer_ms=0.0, post_ms=0.0): - agora = time.time() + def _store_segmentacao( + self, + predictions, + status_result, + analise, + frame_ts, + res, + runner_timing_ms=None, + infer_ms=0.0, + post_ms=0.0, + ): + result_ts = time.time() + frame_ts = float(frame_ts or 0.0) + runner_timing_ms = dict(runner_timing_ms or {}) with self._cache_lock: self._ultimo_predictions = predictions - self._ultimo_seg_aux_result = aux_result + self._ultimo_status_result = status_result self._ultima_analise_segmentacao = analise self._seg_cache = { - "ts": agora, - "frame_ts": frame_ts or 0.0, + "ts": result_ts, + "result_ts": result_ts, + "frame_ts": frame_ts, "predictions": predictions, - "aux_result": aux_result, + "status_result": status_result, "analise": analise, "res": res, "infer_ms": float(infer_ms or 0.0), "post_ms": float(post_ms or 0.0), + "runner_timing_ms": runner_timing_ms, + "pipeline_from_frame_ms": ( + max(0.0, (result_ts - frame_ts) * 1000.0) if frame_ts > 0 else None + ), } self._set_pub_cache("segmentacao", analise.get("dados_visuais", {})) def _store_detection(self, detections, ts, res): - agora = time.time() + result_ts = time.time() + frame_ts = float(ts or 0.0) with self._cache_lock: self._ultimo_detections = detections or [] self._det_cache = { - "ts": agora, - "frame_ts": ts or 0.0, + "ts": result_ts, # compatibilidade + "result_ts": result_ts, + "frame_ts": frame_ts, "detections": detections or [], "res": res, + "pipeline_from_frame_ms": ( + max(0.0, (result_ts - frame_ts) * 1000.0) if frame_ts > 0 else None + ), } self._set_pub_cache("deteccao", detections or []) - def _store_grid(self, snapshot, grid_conf): - agora = time.time() + def _store_grid( + self, + snapshot, + grid_conf, + depth_frame_ts=0.0, + seg_frame_ts=0.0, + det_frame_ts=0.0, + ): + result_ts = time.time() with self._cache_lock: self._ultimo_snapshot = snapshot self._ultima_grid_conf = grid_conf self._grid_cache = { - "ts": agora, + "ts": result_ts, # compatibilidade + "result_ts": result_ts, + "depth_frame_ts": float(depth_frame_ts or 0.0), + "seg_frame_ts": float(seg_frame_ts or 0.0), + "det_frame_ts": float(det_frame_ts or 0.0), "snapshot": snapshot, "grid_conf": grid_conf, } @@ -790,7 +902,10 @@ class CameraManager: # ============================================================ def _processar_grid(self, depth_frame, segmentacao, dados_visuais_seg, deteccoes, velocidade_ms): + t_total0 = time.perf_counter() + try: + t_imu0 = time.perf_counter() imu_roll = 0.0 try: @@ -799,8 +914,13 @@ class CameraManager: except Exception: imu_roll = 0.0 - grid_ref = self.grid_ref_builder.atualizar_por_pitch(imu_roll) + t_imu1 = time.perf_counter() + t_ref0 = time.perf_counter() + grid_ref = self.grid_ref_builder.atualizar_por_pitch(imu_roll) + t_ref1 = time.perf_counter() + + t_builder0 = time.perf_counter() grid_conf = self.grid_builder.build( depth_mm=depth_frame, seg_ids=segmentacao, @@ -808,84 +928,222 @@ class CameraManager: grid_shape=self.grid_ref_builder.grid_shape, deteccoes=deteccoes, ) + t_builder1 = time.perf_counter() if grid_conf is None: - return None, None + return None, None, { + "imu_ms": (t_imu1 - t_imu0) * 1000.0, + "grid_ref_ms": (t_ref1 - t_ref0) * 1000.0, + "builder_ms": (t_builder1 - t_builder0) * 1000.0, + "fuser_ms": 0.0, + "total_ms": (time.perf_counter() - t_total0) * 1000.0, + } + t_fuser0 = time.perf_counter() snapshot = self.costmap_fuser.update( grid_conf, velocidade_ms=velocidade_ms, status_seg=dados_visuais_seg.get("status_corredor"), ) + t_fuser1 = time.perf_counter() - return snapshot, grid_conf + return snapshot, grid_conf, { + "imu_ms": (t_imu1 - t_imu0) * 1000.0, + "grid_ref_ms": (t_ref1 - t_ref0) * 1000.0, + "builder_ms": (t_builder1 - t_builder0) * 1000.0, + "fuser_ms": (t_fuser1 - t_fuser0) * 1000.0, + "total_ms": (time.perf_counter() - t_total0) * 1000.0, + } except Exception as e: self.mostrar_log(f"[visual] erro processando grid: {e}") - return None, None + return None, None, { + "total_ms": (time.perf_counter() - t_total0) * 1000.0, + } # ============================================================ # Preview / stream # ============================================================ - def get_selected_frame(self, frame_type: TipoFrameCamera): + def get_selected_frame( + self, + frame_type: TipoFrameCamera, + aguardar_atualizacao: bool = False, + timeout_s: Optional[float] = None, + ): + """ + Retorna somente frames já produzidos pelo fluxo lateral de preview. + + O runtime crítico nunca renderiza imagem aqui. A chamada apenas solicita + ao produtor de preview que atualize aquele tipo e consome o cache pronto. + + PosProcessamento é a única exceção deliberada: representa o RGB limpo em + resolução original e é obtido diretamente do último frame para salvamento. + """ try: try: tipo_enum = TipoFrameCamera(frame_type) except Exception: tipo_enum = frame_type - # Se visual não está pronto, ainda tenta devolver último frame válido. - if not self._visual_disponivel_para_frame(): - return self._get_cached_preview_frame(tipo_enum) - - with self._cache_lock: - rgb = self._copiar_frame(self._ultimo_rgb_frame) - depth = self._copiar_frame(self._ultimo_depth_frame) - pred = self._ultimo_predictions - dets = list(self._ultimo_detections or []) - snapshot = self._ultimo_snapshot - - params = getattr(self.camera, "parametros", {}) if self.camera is not None else {} - - # PosProcessamento no Visual Worker = imagem RGB original em melhor qualidade. if tipo_enum == TipoFrameCamera.PosProcessamento: - if self._frame_valido(rgb): - self._set_cached_preview_frame(tipo_enum, rgb) - return self._copiar_frame(rgb) + return self._get_latest_rgb_frame_copy() - return self._get_cached_preview_frame(tipo_enum) + cached, cached_ts = self._get_cached_preview_info(tipo_enum, copiar=False) - frame = self.renderer.get_selected_frame( - frame_type=tipo_enum, - rgb_frame=rgb, - pred_ids=pred, - depth_frame=depth, - detections=dets, - snapshot=snapshot, - camera_params=params, - alpha=0.50, + if not self._visual_disponivel_para_frame(): + return self._copiar_frame(cached) + + self._solicitar_preview(tipo_enum) + + agora = time.time() + cache_fresco = ( + self._frame_valido(cached) + and cached_ts > 0.0 + and (agora - cached_ts) <= max(self.preview_cache_fresh_s, 0.0) ) - if self._frame_valido(frame): - self._set_cached_preview_frame(tipo_enum, frame) - return self._copiar_frame(frame) + if not aguardar_atualizacao or cache_fresco: + return self._copiar_frame(cached) - # Se o renderer falhou ou faltou algum insumo naquele instante, - # devolve o último frame válido desse mesmo tipo. - cached = self._get_cached_preview_frame(tipo_enum) - if self._frame_valido(cached): - return cached + timeout = self.preview_request_timeout_s if timeout_s is None else float(timeout_s) + deadline = time.time() + max(timeout, 0.0) - return None + with self._preview_cond: + while time.time() < deadline: + atual_ts = float( + self._ultimo_frame_preview_ts_por_tipo.get( + self._frame_cache_key(tipo_enum), 0.0 + ) or 0.0 + ) + if atual_ts > cached_ts: + break + + restante = deadline - time.time() + if restante <= 0: + break + self._preview_cond.wait(timeout=min(restante, 0.05)) + + return self._get_cached_preview_frame(tipo_enum) except Exception as e: self.mostrar_log(f"[visual] erro em get_selected_frame({frame_type}): {e}") + return self._get_cached_preview_frame(frame_type) + def _solicitar_preview(self, frame_type) -> bool: + try: + tipo_enum = TipoFrameCamera(frame_type) + except Exception: + tipo_enum = frame_type + + if tipo_enum == TipoFrameCamera.PosProcessamento: + return False + + if tipo_enum not in self.FRAME_TYPES_PREVIEW: + return False + + key = self._frame_cache_key(tipo_enum) + + with self._preview_cond: + self._preview_requests.add(key) + self._preview_cond.notify_all() + + return True + + def _consumir_solicitacoes_preview(self): + with self._preview_cond: + if not self._preview_requests: + return [] + + keys = tuple(self._preview_requests) + self._preview_requests.clear() + + tipos = [] + for key in keys: try: - return self._get_cached_preview_frame(frame_type) + tipos.append(TipoFrameCamera[key]) except Exception: - return None + continue + return tipos + + def _snapshot_preview_inputs(self, frame_type): + """Fotografa apenas referências necessárias, mantendo o lock crítico curto.""" + try: + tipo_enum = TipoFrameCamera(frame_type) + except Exception: + tipo_enum = frame_type + + rgb = None + depth = None + pred = None + dets = None + snapshot = None + + precisa_rgb = tipo_enum in { + TipoFrameCamera.Rgb, + TipoFrameCamera.Overlay, + TipoFrameCamera.MatrizCusto, + TipoFrameCamera.Deteccoes, + TipoFrameCamera.Debug, + } + precisa_depth = tipo_enum == TipoFrameCamera.Heatmap + precisa_pred = tipo_enum in { + TipoFrameCamera.Segmentacao, + TipoFrameCamera.Overlay, + TipoFrameCamera.Debug, + } + precisa_dets = tipo_enum == TipoFrameCamera.Deteccoes + precisa_snapshot = tipo_enum in { + TipoFrameCamera.MatrizCusto, + TipoFrameCamera.Debug, + } + + with self._cache_lock: + if precisa_rgb: + rgb = self._ultimo_rgb_frame + if precisa_depth: + depth = self._ultimo_depth_frame + if precisa_pred: + pred = self._ultimo_predictions + if precisa_dets: + dets = list(self._ultimo_detections or []) + if precisa_snapshot: + snapshot = self._ultimo_snapshot + + params = getattr(self.camera, "parametros", {}) if self.camera is not None else {} + + return { + "rgb": rgb, + "depth": depth, + "pred": pred, + "dets": dets, + "snapshot": snapshot, + "params": params, + } + + def _render_preview_frame(self, frame_type): + if self.renderer is None: + return None + + data = self._snapshot_preview_inputs(frame_type) + + return self.renderer.get_selected_frame( + frame_type=frame_type, + rgb_frame=data["rgb"], + pred_ids=data["pred"], + depth_frame=data["depth"], + detections=data["dets"], + snapshot=data["snapshot"], + camera_params=data["params"], + alpha=0.50, + ) + + def _get_latest_rgb_frame_copy(self): + with self._cache_lock: + frame = self._ultimo_rgb_frame + + # A cópia pesada acontece fora do lock do runtime. + return self._copiar_frame(frame) def _frame_valido(self, frame) -> bool: return ( @@ -914,25 +1172,32 @@ class CameraManager: return False key = self._frame_cache_key(frame_type) - ts = time.time() if ts is None else ts - frame_copy = self._copiar_frame(frame) + ts = time.time() if ts is None else float(ts) - if not self._frame_valido(frame_copy): - return False - - with self._cache_lock: - self._ultimo_frame_preview_por_tipo[key] = frame_copy + # O frame foi criado pelo produtor de preview e não volta a ser mutado. + # Portanto o cache pode guardar a própria referência sem uma cópia extra. + with self._preview_cond: + self._ultimo_frame_preview_por_tipo[key] = frame self._ultimo_frame_preview_ts_por_tipo[key] = ts + self._preview_cond.notify_all() return True - def _get_cached_preview_frame(self, frame_type): + def _get_cached_preview_frame(self, frame_type, copiar=True): + frame, _ = self._get_cached_preview_info(frame_type, copiar=False) + return self._copiar_frame(frame) if copiar else frame + + def _get_cached_preview_info(self, frame_type, copiar=True): key = self._frame_cache_key(frame_type) - with self._cache_lock: + with self._preview_lock: frame = self._ultimo_frame_preview_por_tipo.get(key) + ts = float(self._ultimo_frame_preview_ts_por_tipo.get(key, 0.0) or 0.0) - return self._copiar_frame(frame) + if copiar: + frame = self._copiar_frame(frame) + + return frame, ts def _normalizar_frame_para_uint8(self, frame): """ @@ -967,6 +1232,70 @@ class CameraManager: # Loops # ============================================================ + def _iniciar_loop_preview(self, freq=2.0): + """ + Produtor lateral de preview. + + Dorme sem custo quando não há consumidores. Quando recebe solicitações, + renderiza somente os tipos pedidos e sempre fora do lock do runtime. + """ + def loop(): + ultimo_render_por_tipo = {} + + while True: + try: + if self.camera is None or not self.operante: + with self._preview_cond: + self._preview_cond.wait(timeout=0.25) + continue + + with self._preview_cond: + if not self._preview_requests: + self._preview_cond.wait(timeout=0.50) + + tipos = self._consumir_solicitacoes_preview() + if not tipos: + continue + + periodo_min = 1.0 / max(float(self.preview_fps or freq), 0.1) + + for tipo in tipos: + key = self._frame_cache_key(tipo) + agora = time.time() + ultimo = float(ultimo_render_por_tipo.get(key, 0.0) or 0.0) + + # Limita o custo máximo do preview mesmo se o operador/UI + # fizer polling acima da frequência configurada. + if ultimo > 0.0 and (agora - ultimo) < periodo_min: + continue + + t0 = time.perf_counter() + frame = self._render_preview_frame(tipo) + t1 = time.perf_counter() + + if not self._frame_valido(frame): + self.perf.inc("preview_sem_frame") + continue + + t_cache0 = time.perf_counter() + self._set_cached_preview_frame(tipo, frame, ts=time.time()) + t_cache1 = time.perf_counter() + + ultimo_render_por_tipo[key] = time.time() + + self.perf.tick( + "preview", + latencia_ms=(t_cache1 - t0) * 1000.0, + render_ms=(t1 - t0) * 1000.0, + cache_write_ms=(t_cache1 - t_cache0) * 1000.0, + ) + + except Exception as e: + self.mostrar_log(f"[visual] erro no produtor de preview: {e}") + time.sleep(0.05) + + threading.Thread(target=loop, daemon=True).start() + def _iniciar_loop_segmentacao(self, freq=8.0): def loop(): ultimo_frame_ts = 0.0 @@ -992,7 +1321,7 @@ class CameraManager: t_get0 = time.perf_counter() rgb_frame, frame_ts, res = self.get_rgb_frame() - rgb_frame = self._validar_frame_rgb_para_segmentacao(rgb_frame) + rgb_frame = self._validar_frame_camera_para_segmentacao(rgb_frame) t_get1 = time.perf_counter() if rgb_frame is None or frame_ts is None: @@ -1008,23 +1337,26 @@ class CameraManager: ultimo_frame_ts = frame_ts t_inf0 = time.perf_counter() - predictions, ts, roi_resized, roi_info, aux_result = self.seg_runner.infer_ids(rgb_frame) + infer_result = self.seg_runner.infer(rgb_frame) t_inf1 = time.perf_counter() - if predictions is None or ts is None: - self.perf.inc("segmentacao_pred_none") - continue + predictions = infer_result.seg_ids + status_result = infer_result.status t_post0 = time.perf_counter() - analise = self.segmentacao_manager.analisar(predictions, aux_result=aux_result) + analise = self.segmentacao_manager.analisar( + predictions, + status_result=status_result, + ) t_post1 = time.perf_counter() self._store_segmentacao( predictions=predictions, - aux_result=aux_result, + status_result=status_result, analise=analise, frame_ts=frame_ts, res=res, + runner_timing_ms=infer_result.timing_ms, infer_ms=(t_inf1 - t_inf0) * 1000.0, post_ms=(t_post1 - t_post0) * 1000.0, ) @@ -1036,6 +1368,13 @@ class CameraManager: latencia_ms=(t_loop1 - t_loop0) * 1000.0, get_rgb_ms=(t_get1 - t_get0) * 1000.0, infer_ms=(t_inf1 - t_inf0) * 1000.0, + geometry_ms=float(infer_result.timing_ms.get("prepare_total_ms", 0.0)), + crop_ms=float(infer_result.timing_ms.get("crop_ms", 0.0)), + resize_ms=float(infer_result.timing_ms.get("resize_ms", 0.0)), + color_ms=float(infer_result.timing_ms.get("color_ms", 0.0)), + session_ms=float(infer_result.timing_ms.get("session_ms", 0.0)), + decode_ms=float(infer_result.timing_ms.get("decode_ms", 0.0)), + runner_total_ms=float(infer_result.timing_ms.get("total_ms", 0.0)), post_ms=(t_post1 - t_post0) * 1000.0, redis_ms=0.0, frame_ts=frame_ts, @@ -1110,8 +1449,8 @@ class CameraManager: def _iniciar_loop_grid(self, freq=5.0): def loop(): ultimo_depth_ts = 0.0 - ultimo_seg_ts = 0.0 - ultimo_det_ts = 0.0 + ultimo_seg_result_ts = 0.0 + ultimo_det_result_ts = 0.0 while True: t_wall0 = time.time() @@ -1151,8 +1490,13 @@ class CameraManager: dados_visuais = analise_seg.get("dados_visuais", {}) deteccoes = list(det_cache.get("detections", []) or []) - seg_ts = float(seg_cache.get("ts", 0.0) or 0.0) - det_ts = float(det_cache.get("ts", 0.0) or 0.0) + # Há dois relógios diferentes e ambos são úteis: + # frame_ts = quando o dado nasceu na câmera; + # result_ts = quando o processamento correspondente terminou. + seg_frame_ts = float(seg_cache.get("frame_ts", 0.0) or 0.0) + det_frame_ts = float(det_cache.get("frame_ts", 0.0) or 0.0) + seg_result_ts = float(seg_cache.get("result_ts", seg_cache.get("ts", 0.0)) or 0.0) + det_result_ts = float(det_cache.get("result_ts", det_cache.get("ts", 0.0)) or 0.0) t_cache1 = time.perf_counter() if segmentacao is None: @@ -1160,14 +1504,18 @@ class CameraManager: time.sleep(0.02) continue - if depth_ts == ultimo_depth_ts and seg_ts == ultimo_seg_ts and det_ts == ultimo_det_ts: + if ( + depth_ts == ultimo_depth_ts + and seg_result_ts == ultimo_seg_result_ts + and det_result_ts == ultimo_det_result_ts + ): self.perf.inc("grid_pacote_repetido") time.sleep(0.005) continue ultimo_depth_ts = depth_ts - ultimo_seg_ts = seg_ts - ultimo_det_ts = det_ts + ultimo_seg_result_ts = seg_result_ts + ultimo_det_result_ts = det_result_ts try: vel = get_velocidade_atual_ms() @@ -1175,7 +1523,7 @@ class CameraManager: vel = 0.0 t_grid0 = time.perf_counter() - snapshot, grid_conf = self._processar_grid( + snapshot, grid_conf, grid_timing = self._processar_grid( depth_frame=depth_frame, segmentacao=segmentacao, dados_visuais_seg=dados_visuais, @@ -1189,7 +1537,15 @@ class CameraManager: time.sleep(0.005) continue - self._store_grid(snapshot=snapshot, grid_conf=grid_conf) + self._store_grid( + snapshot=snapshot, + grid_conf=grid_conf, + depth_frame_ts=depth_ts, + seg_frame_ts=seg_frame_ts, + det_frame_ts=det_frame_ts, + ) + + agora = time.time() self.perf.tick( "grid", @@ -1197,16 +1553,23 @@ class CameraManager: get_depth_ms=(t_depth1 - t_depth0) * 1000.0, cache_read_ms=(t_cache1 - t_cache0) * 1000.0, build_grid_ms=(t_grid1 - t_grid0) * 1000.0, - fuser_ms=0.0, + imu_ms=float(grid_timing.get("imu_ms", 0.0) or 0.0), + grid_ref_ms=float(grid_timing.get("grid_ref_ms", 0.0) or 0.0), + grid_builder_ms=float(grid_timing.get("builder_ms", 0.0) or 0.0), + fuser_ms=float(grid_timing.get("fuser_ms", 0.0) or 0.0), redis_ms=0.0, depth_ts=depth_ts, - seg_ts=seg_ts, - det_ts=det_ts, - idade_depth_ms=(time.time() - depth_ts) * 1000.0 if depth_ts else None, - idade_seg_ms=(time.time() - seg_ts) * 1000.0 if seg_ts else None, - idade_det_ms=(time.time() - det_ts) * 1000.0 if det_ts else None, - sync_depth_seg_ms=abs(depth_ts - seg_ts) * 1000.0 if depth_ts and seg_ts else None, - sync_depth_det_ms=abs(depth_ts - det_ts) * 1000.0 if depth_ts and det_ts else None, + seg_ts=seg_frame_ts, + det_ts=det_frame_ts, + seg_result_ts=seg_result_ts, + det_result_ts=det_result_ts, + idade_depth_frame_ms=(agora - depth_ts) * 1000.0 if depth_ts else None, + idade_seg_frame_ms=(agora - seg_frame_ts) * 1000.0 if seg_frame_ts else None, + idade_det_frame_ms=(agora - det_frame_ts) * 1000.0 if det_frame_ts else None, + idade_seg_resultado_ms=(agora - seg_result_ts) * 1000.0 if seg_result_ts else None, + idade_det_resultado_ms=(agora - det_result_ts) * 1000.0 if det_result_ts else None, + sync_depth_seg_frame_ms=(abs(depth_ts - seg_frame_ts) * 1000.0) if depth_ts and seg_frame_ts else None, + sync_depth_det_frame_ms=(abs(depth_ts - det_frame_ts) * 1000.0) if depth_ts and det_frame_ts else None, n_dets=len(deteccoes), ) @@ -1306,13 +1669,15 @@ class CameraManager: camera_ctx.get("frame_type", TipoFrameCamera.Rgb.value) ) - frame = self.get_selected_frame(frame_type) + self._solicitar_preview(frame_type) + frame = self._get_cached_preview_frame(frame_type, copiar=False) if frame is not None: self.camera.enviar_frame_tcp(frame) if self.debug_visual and self._visual_disponivel_para_frame(): - dbg = self.get_selected_frame(TipoFrameCamera.Debug) + self._solicitar_preview(TipoFrameCamera.Debug) + dbg = self._get_cached_preview_frame(TipoFrameCamera.Debug, copiar=False) if dbg is not None: cv2.imshow("Visual Worker Debug", dbg) cv2.waitKey(1) @@ -1366,7 +1731,8 @@ class CameraManager: if ctrl is not None: ctrl.update_from_redis() - self._publicar_e_logar_performance() + if self.telemetry_enabled: + self._publicar_e_logar_performance() except Exception as e: self.mostrar_log(f"Erro no loop supervisor visual: {e}") @@ -1417,11 +1783,32 @@ class CameraManager: self.seg_config = cfg self.debug_visual = bool(cfg.get("debug_visual", self.debug_visual)) - self.debug_perf = bool(cfg.get("debug_perf", self.debug_perf)) + + telemetry_cfg = cfg.get("telemetry", {}) or {} + self.telemetry_enabled = bool(telemetry_cfg.get("enabled", self.telemetry_enabled)) + self.telemetry_console = bool(telemetry_cfg.get("console_performance", self.telemetry_console)) + self.telemetry_publish_runtime = bool(telemetry_cfg.get("publish_runtime_info", self.telemetry_publish_runtime)) + self.telemetry_runner_timing = bool(telemetry_cfg.get("runner_detailed_timing", self.telemetry_runner_timing)) + + include_probs = bool(telemetry_cfg.get("include_status_probs", False)) + if self.seg_runner is not None: + self.seg_runner.include_status_probs = include_probs + if self.segmentacao_manager is not None: + self.segmentacao_manager.config.include_model_probs = include_probs self.posproc_intervalo_min_s = float( cfg.get("posproc_intervalo_min_s", self.posproc_intervalo_min_s) ) + self.preview_fps = float(cfg.get("preview_fps", self.preview_fps)) + self.preview_request_timeout_s = float( + cfg.get("preview_request_timeout_s", self.preview_request_timeout_s) + ) + self.preview_cache_fresh_s = float( + cfg.get("preview_cache_fresh_s", self.preview_cache_fresh_s) + ) + self.replay_jpeg_quality = int( + cfg.get("replay_jpeg_quality", self.replay_jpeg_quality) + ) self._ultimo_config_update_ts = agora @@ -1489,19 +1876,38 @@ class CameraManager: if self.camera is not None and hasattr(self.camera, "get_cache_stats"): resumo["camera_cache"] = self.camera.get_cache_stats() + agora = time.time() + seg_frame_ts = float(self._seg_cache.get("frame_ts", 0.0) or 0.0) + seg_result_ts = float(self._seg_cache.get("result_ts", self._seg_cache.get("ts", 0.0)) or 0.0) + det_frame_ts = float(self._det_cache.get("frame_ts", 0.0) or 0.0) + det_result_ts = float(self._det_cache.get("result_ts", self._det_cache.get("ts", 0.0)) or 0.0) + grid_result_ts = float(self._grid_cache.get("result_ts", self._grid_cache.get("ts", 0.0)) or 0.0) + resumo["visual_cache"] = { - "seg_ts": self._seg_cache.get("ts", 0.0), - "det_ts": self._det_cache.get("ts", 0.0), - "grid_ts": self._grid_cache.get("ts", 0.0), + "seg_frame_ts": seg_frame_ts, + "seg_result_ts": seg_result_ts, + "det_frame_ts": det_frame_ts, + "det_result_ts": det_result_ts, + "grid_result_ts": grid_result_ts, + "seg_frame_age_ms": (agora - seg_frame_ts) * 1000.0 if seg_frame_ts else None, + "seg_result_age_ms": (agora - seg_result_ts) * 1000.0 if seg_result_ts else None, + "det_frame_age_ms": (agora - det_frame_ts) * 1000.0 if det_frame_ts else None, + "det_result_age_ms": (agora - det_result_ts) * 1000.0 if det_result_ts else None, + "grid_result_age_ms": (agora - grid_result_ts) * 1000.0 if grid_result_ts else None, + "seg_pipeline_from_frame_ms": self._seg_cache.get("pipeline_from_frame_ms"), + "det_pipeline_from_frame_ms": self._det_cache.get("pipeline_from_frame_ms"), "tem_seg": self._ultimo_predictions is not None, "tem_grid": self._ultimo_snapshot is not None, } resumo["pub_debug"] = self._ultimo_pub_debug + if self.telemetry_publish_runtime and self.seg_runner is not None: + resumo["seg_runtime"] = self.seg_runner.runtime_info() + self._set_pub_cache("performance_visual", resumo) - if self.debug_perf: + if self.telemetry_console: self._logar_performance(resumo) def _logar_performance(self, resumo): @@ -1513,13 +1919,14 @@ class CameraManager: seg = loops.get("segmentacao", {}) det = loops.get("deteccao", {}) grid = loops.get("grid", {}) + preview = loops.get("preview", {}) pub = loops.get("publicacao", {}) self.mostrar_log( "[PERF_VISUAL] " f"fps rgb={self._fps(rgb):.1f} depth={self._fps(depth):.1f} " f"seg={self._fps(seg):.1f} det={self._fps(det):.1f} grid={self._fps(grid):.1f} " - f"pub={self._fps(pub):.1f} | " + f"preview={self._fps(preview):.1f} pub={self._fps(pub):.1f} | " f"period seg={self._fmt(self._per(seg))}ms grid={self._fmt(self._per(grid))}ms | " f"sync={self._fmt(sync.get('rgb_depth_dt_ms'))}ms" ) @@ -1528,17 +1935,46 @@ class CameraManager: "[PERF_DETAIL] " f"SEG total={self._fmt(self._lat(seg))} " f"get={self._fmt(self._m(seg, 'get_rgb_ms'))} " - f"infer={self._fmt(self._m(seg, 'infer_ms'))} " + f"geom={self._fmt(self._m(seg, 'geometry_ms'))} " + f"session={self._fmt(self._m(seg, 'session_ms'))} " + f"decode={self._fmt(self._m(seg, 'decode_ms'))} " f"post={self._fmt(self._m(seg, 'post_ms'))} | " f"GRID total={self._fmt(self._lat(grid))} " f"depth={self._fmt(self._m(grid, 'get_depth_ms'))} " - f"build={self._fmt(self._m(grid, 'build_grid_ms'))} | " + f"imu={self._fmt(self._m(grid, 'imu_ms'))} " + f"ref={self._fmt(self._m(grid, 'grid_ref_ms'))} " + f"builder={self._fmt(self._m(grid, 'grid_builder_ms'))} " + f"fuser={self._fmt(self._m(grid, 'fuser_ms'))} | " f"DET total={self._fmt(self._lat(det))} " f"read={self._fmt(self._m(det, 'oak_read_ms'))} | " + f"PREVIEW total={self._fmt(self._lat(preview))} " + f"render={self._fmt(self._m(preview, 'render_ms'))} | " f"PUB total={self._fmt(self._lat(pub))} " f"redis={self._fmt(self._m(pub, 'redis_ms'))}" ) + if self.telemetry_runner_timing: + self.mostrar_log( + "[PERF_SEG_RUNNER] " + f"crop={self._fmt(self._m(seg, 'crop_ms'))}ms " + f"resize={self._fmt(self._m(seg, 'resize_ms'))}ms " + f"color={self._fmt(self._m(seg, 'color_ms'))}ms " + f"session={self._fmt(self._m(seg, 'session_ms'))}ms " + f"decode={self._fmt(self._m(seg, 'decode_ms'))}ms " + f"runner={self._fmt(self._m(seg, 'runner_total_ms'))}ms" + ) + + self.mostrar_log( + "[PERF_SYNC] " + f"depth-seg-frame={self._fmt(self._m(grid, 'sync_depth_seg_frame_ms'))}ms " + f"depth-det-frame={self._fmt(self._m(grid, 'sync_depth_det_frame_ms'))}ms | " + f"age depth={self._fmt(self._m(grid, 'idade_depth_frame_ms'))}ms " + f"seg_frame={self._fmt(self._m(grid, 'idade_seg_frame_ms'))}ms " + f"seg_result={self._fmt(self._m(grid, 'idade_seg_resultado_ms'))}ms " + f"det_frame={self._fmt(self._m(grid, 'idade_det_frame_ms'))}ms " + f"det_result={self._fmt(self._m(grid, 'idade_det_resultado_ms'))}ms" + ) + def _sleep_loop_adaptativo(self, nome_tarefa, t0_wall, freq_fallback): try: ctrl = self.gpu_controller @@ -1557,17 +1993,13 @@ class CameraManager: def salvar_frames(self, tipos: list, nome: str, pasta="frames_salvos"): """ - Salva frames sob demanda. + Salva frames sob demanda sem renderizar no chamador. - Padrão: - - Frames de replay/debug: - JPEG leve, qualidade configurada em replay_jpeg_quality. - - - Tipo PosProcessamento: - imagem original em qualidade máxima, sem compressão destrutiva forte. - No Visual Worker, isso é apenas a imagem RGB limpa. + - Replay/debug: usa somente o preview já produzido pelo Preview Producer + e salva JPEG leve. + - PosProcessamento: única exceção de alta qualidade, usando RGB original + e PNG em thread assíncrona. """ - if not self._visual_disponivel_para_frame(): return [] @@ -1577,8 +2009,6 @@ class CameraManager: frame_name = nome.strip() if nome else datetime.datetime.now().strftime("%Y%m%d_%H%M%S_%f")[:-3] - replay_jpeg_quality = 70 - for frame_type in tipos: try: tipo_enum = TipoFrameCamera(frame_type) @@ -1586,26 +2016,11 @@ class CameraManager: self.mostrar_log(f"[visual] TipoFrameCamera inválido para salvar: {frame_type}") continue - frame = self.get_selected_frame(tipo_enum) - - if not self._frame_valido(frame): - self.mostrar_log( - f"⚠️ Frame visual não salvo | " - f"tipo={tipo_enum.name} " - f"nome={frame_name}_{tipo_enum.name} " - f"motivo=sem_frame_valido_em_cache" - ) - continue - - frame = self._copiar_frame(frame) - nome_base = f"{frame_name}_{tipo_enum.name}" # ==================================================== - # Pós-processamento + # Pós-processamento: RGB original / alta qualidade # ==================================================== - # No Visual Worker, o dado científico é só a imagem limpa. - # Salva com qualidade máxima. if tipo_enum == TipoFrameCamera.PosProcessamento: pode_salvar, motivo = self._pode_salvar_posprocessamento() @@ -1617,66 +2032,76 @@ class CameraManager: ) continue - frame_pp = frame.copy() + frame_pp = self._get_latest_rgb_frame_copy() + if not self._frame_valido(frame_pp): + self._finalizar_salvamento_posprocessamento() + self.mostrar_log( + f"⚠️ Frame visual não salvo | tipo={tipo_enum.name} " + f"nome={nome_base} motivo=sem_rgb_original" + ) + continue def salvar_pp_async(nome_base_async, pasta_async, frame_async): try: t_pp0 = time.perf_counter() - caminho = os.path.join(pasta_async, f"{nome_base_async}.png") - ok = cv2.imwrite(caminho, frame_async) ok_final = bool(ok and os.path.exists(caminho)) - t_pp1 = time.perf_counter() self.mostrar_log( - f"[visual][SAVE_PP] salvo async | " - f"nome={nome_base_async} " - f"ok={ok_final} " - f"tempo_ms={(t_pp1 - t_pp0) * 1000.0:.1f}" + f"[visual][SAVE_PP] salvo async | nome={nome_base_async} " + f"ok={ok_final} tempo_ms={(t_pp1 - t_pp0) * 1000.0:.1f}" ) - except Exception as e: self.mostrar_log(f"❌ Erro no visual SAVE_PP async: {e}") - finally: self._finalizar_salvamento_posprocessamento() threading.Thread( target=salvar_pp_async, args=(nome_base, pasta, frame_pp), - daemon=True + daemon=True, ).start() frames_salvos.append(f"ASYNC_STARTED:{nome_base}") continue # ==================================================== - # Replay/debug + # Replay/debug: somente preview pronto/cacheado # ==================================================== - # Salva JPEG leve para não entupir disco durante operação. - caminho = os.path.join(pasta, f"{nome_base}.jpg") + frame = self.get_selected_frame( + tipo_enum, + aguardar_atualizacao=True, + timeout_s=self.preview_request_timeout_s, + ) - if getattr(frame, "dtype", None) != np.uint8: - frame_salvar = self._normalizar_frame_para_uint8(frame) - else: - frame_salvar = frame + if not self._frame_valido(frame): + self.mostrar_log( + f"⚠️ Frame visual não salvo | tipo={tipo_enum.name} " + f"nome={nome_base} motivo=sem_preview_valido" + ) + continue + + caminho = os.path.join(pasta, f"{nome_base}.jpg") + frame_salvar = ( + self._normalizar_frame_para_uint8(frame) + if getattr(frame, "dtype", None) != np.uint8 + else frame + ) ok = cv2.imwrite( caminho, frame_salvar, - [int(cv2.IMWRITE_JPEG_QUALITY), int(replay_jpeg_quality)] + [int(cv2.IMWRITE_JPEG_QUALITY), int(self.replay_jpeg_quality)], ) if ok and os.path.exists(caminho): frames_salvos.append(caminho) else: self.mostrar_log( - f"❌ Falha ao salvar frame visual | " - f"tipo={tipo_enum.name} " - f"nome={nome_base} " - f"caminho={caminho}" + f"❌ Falha ao salvar frame visual | tipo={tipo_enum.name} " + f"nome={nome_base} caminho={caminho}" ) return frames_salvos diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/config.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/config.py index cacf2a798..cf9f8f9d7 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/config.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/config.py @@ -50,7 +50,7 @@ _CONFIG_LOCK = threading.Lock() # ============================================================ -# CONFIG V1 - VISUAL WORKER +# CONFIG V2 - VISUAL WORKER / ONNX FIELD CONTRACT # ============================================================ @@ -117,12 +117,39 @@ VISUAL_DEFAULT_CONFIG = { "debug_visual": False, "posproc_intervalo_min_s": 5.0, - # Publica logs de performance no console. - "debug_perf": False, - # Tamanho dos frames de preview/debug enviados ao C#. "preview_size": (1280, 720), + # Preview producer lateral: só trabalha quando há consumidor. + "preview_fps": 2.0, + "preview_request_timeout_s": 0.35, + "preview_cache_fresh_s": 0.75, + + # Replay/debug usa JPEG leve. PosProcessamento continua em PNG/RGB original. + "replay_jpeg_quality": 70, + + # Telemetria técnica. Tudo que custa payload/log extra fica atrás deste bloco. + "telemetry": { + # Publica performance_visual no Redis. + "enabled": True, + + # Imprime resumo periódico no console. + "console_performance": False, + + # Abre o timing da inferência em crop/resize/cor/session/decode. + "runner_detailed_timing": True, + + # Inclui o vetor completo de 8 probabilidades do status no debug. + # Desligado por padrão porque top1/top2/margin/entropy já bastam no campo. + "include_status_probs": False, + + # Anexa contrato/runtime do ONNX em performance_visual. + "publish_runtime_info": True, + + # Loga contrato/provider/cache na inicialização. + "log_runtime_startup": True, + }, + # ============================================================ # 2) Frequências dos loops @@ -144,100 +171,81 @@ VISUAL_DEFAULT_CONFIG = { # ============================================================ - # 3) Modelo ONNX/TensorRT - segmentação de ruas/corredor + # 3) Runtime oficial do modelo de corredor - ONNX FIELD V2 # ============================================================ - # Runtime oficial da v1. - "runtime_backend": "onnx", - "onnx_provider": "tensorrt", + # O caminho do ONNX vem do Redis/equipamento: path_ia_model_ruas_seg. + # Resolução, ROI, mean/std, layout, nomes de I/O e classes NÃO ficam aqui. + # Tudo isso é lido da metadata embutida no próprio ONNX. + "model_runtime": { + "provider": "tensorrt", - # Resolução esperada pelo ONNX: [W, H]. - "ia_resolution": [1024, 576], + # Se True: TensorRT -> CUDA -> CPU, sempre expondo degradação na telemetria. + # Se False: falha a inicialização quando o provider solicitado não estiver disponível. + "allow_provider_fallback": True, - # ROI vertical da segmentação. - # 0.0 + 1.0 = frame inteiro. - "ia_roi_begin": 0.0, - "ia_roi_size": 1.0, + # CameraOak configura ColorCamera.ColorOrder.RGB. O ONNX também exige RGB, + # então o hot path evita conversão de cor. Se a fonte mudar no futuro, + # BGR continua suportado explicitamente pelo runner. + "camera_frame_color": "RGB", - # Contrato validado do modelo modelseg-2_0.onnx: - # entrada: float32 [1, 3, 576, 1024], RGB 0..1, NCHW - # saída 1: semantic_logits - # saída 2: label_probs - "onnx_has_preprocess": False, - "onnx_input_layout": "nchw_float", - "onnx_input_scale": 1.0 / 255.0, - "onnx_seg_output_format": "logits", + # TensorRT. O cache real vira // automaticamente. + "trt_fp16_enable": True, + "trt_engine_cache_enable": True, + "trt_engine_cache_root": "./trt_cache_visual_worker", + "trt_max_workspace_size": None, - # Nomes das entradas/saídas. - # input_name None deixa o runner detectar automaticamente. - "onnx_input_name": None, - "onnx_seg_output_name": "semantic_logits", - "onnx_aux_output_name": "label_probs", - "onnx_aux_output_format": "probs", + # ONNX Runtime. + "ort_intra_op_num_threads": 1, + "ort_inter_op_num_threads": 1, - # Auxiliar do status de corredor. - # True = retorna label_id, label_name, label_conf. - # False = inclui também vetor de probabilidades. - "use_compact_aux": True, + # Se existir .contract.json, valida o SHA256 do ONNX no startup. + "verify_sidecar_sha256": True, - # Logs internos do runner ONNX. - "debug_timing": False, - "debug_session": True, + # Checagem de IDs únicos a cada inferência. Útil em laboratório, desnecessária no campo. + "validate_outputs_each_inference": False, + }, # ============================================================ - # 4) TensorRT + # 4) Análise geométrica + resolução temporal do status # ============================================================ - "trt_fp16_enable": True, - "trt_engine_cache_enable": True, - "trt_engine_cache_path": "./trt_cache_visual_worker", - - # Deixe None salvo se não quiser fixar workspace. - "trt_max_workspace_size": None, - - # Threads do ONNX Runtime. - # Para TensorRT/CUDA, 1 costuma ser suficiente e evita ruído. - "onnx_intra_op_num_threads": 1, - "onnx_inter_op_num_threads": 1, - - - # ============================================================ - # 5) SegmentacaoManager v1 - # ============================================================ - # Este bloco controla apenas a análise da máscara: - # pred_ids + label_probs -> dados_visuais. + # IDs semânticos e nomes de status vêm do contrato ONNX. "segmentacao": { - # Inclui debug textual no payload de segmentação. - # Não gera imagem. "include_debug": False, - - # Inclui timing interno do SegmentacaoManager. "debug_timing": False, - # ID das classes no modelo de segmentação. - "id_nao_navegavel": 0, - "id_navegavel": 1, - - # Frações verticais usadas para estimar centro/ângulo do corredor. - # 0.0 = topo, 1.0 = base. + # Frações verticais para centro/largura/ângulo. 0=topo, 1=base. "scanline_fracs": (0.96, 0.86, 0.74, 0.62, 0.50, 0.38, 0.26), - - # A região próxima ao robô está na base da imagem. "near_is_bottom": True, - # Suavização das saídas usadas pelo controle. + # Suavização geométrica. "ema_alpha_ang": 0.25, "ema_alpha_lat": 0.25, "ema_alpha_conf": 0.20, - # Histerese temporal do status do corredor. - "status_window_s": 1.5, + # Status v2: voto temporal ponderado + troca rápida quando a cabeça está muito segura. + "status_window_s": 0.90, + "status_decay_tau_s": 0.35, "status_expected_fps": 10.0, - # Confiança da cabeça auxiliar ONNX. - # >= accept: modelo manda. - # >= soft: modelo ajuda quando heurística está indefinida. - "model_conf_accept": 0.70, - "model_conf_soft": 0.45, + # Modelo manda quando confiança E margem são boas. + "model_conf_accept": 0.72, + "model_margin_accept": 0.15, + + # Zona intermediária: só ajuda quando concorda com heurística ou ela está indefinida. + "model_conf_soft": 0.55, + "model_margin_soft": 0.06, + + # Transições muito fortes podem furar a inércia após N frames consecutivos. + "model_conf_fast": 0.90, + "model_margin_fast": 0.30, + "fast_switch_min_frames": 2, + + # Pesos base do voto temporal. + "status_weight_model": 1.35, + "status_weight_model_soft": 1.00, + "status_weight_heuristic": 0.75, + "status_weight_agreement_bonus": 0.40, # Grid leve interna para centro/fallback do corredor. "corridor_grid_rows": 6, @@ -250,7 +258,7 @@ VISUAL_DEFAULT_CONFIG = { "prefer_prev_weight": 0.65, "prefer_width_weight": 1.00, - # Heurística de status quando o modelo auxiliar não está confiante. + # Heurística fica como fallback/segundo sensor até termos evidência de campo para removê-la. "thr_parado_global": 0.30, "thr_direcionando_global": 0.82, "thr_caminhando_score": 0.48, @@ -260,7 +268,7 @@ VISUAL_DEFAULT_CONFIG = { # ============================================================ - # 6) Grid de confiança/custo + # 5) Grid de confiança/custo # ============================================================ "grid": { # Formato global da grid: (cols, rows). @@ -408,7 +416,7 @@ VISUAL_DEFAULT_CONFIG = { # ============================================================ - # 7) GPU Priority Controller + # 6) GPU Priority Controller # ============================================================ # Controla FPS do Visual Worker conforme saúde do Weed Worker. "gpu_priority": { @@ -510,24 +518,16 @@ DET_DEFAULT_CONFIG = { def aplicar_overrides_redis_seg(cfg: dict) -> dict: equipamento = ContextoGlobalRedis.get_equipamento() - # Modelo ONNX oficial de ruas/corredor. - cfg["ia_onnx_path"] = equipamento.get("path_ia_model_ruas_seg") - - # Labelmap da segmentação. - cfg["ia_labelmap_path"] = equipamento.get("path_ia_labelmap_ruas_seg") + # Único artefato de IA exigido pelo Visual Worker v2. + # Classes, normalização, ROI e I/O vivem dentro do contrato do ONNX. + cfg["model_path"] = equipamento.get("path_ia_model_ruas_seg") return cfg def normalizar_config_runtime_seg(cfg: dict) -> dict: - # Alias único interno, para logs e validações. - cfg["onnx_model_path"] = cfg.get("ia_onnx_path") - - if not cfg.get("ia_onnx_path"): - mostrar_log("[WARN] path do modelo ONNX de ruas não definido no Redis/equipamento.") - - if not cfg.get("ia_labelmap_path"): - mostrar_log("[WARN] path do labelmap de ruas não definido no Redis/equipamento.") + if not cfg.get("model_path"): + mostrar_log("[WARN] path do ONNX de corredor não definido no Redis/equipamento.") # Garante tuplas onde as dataclasses esperam tuplas. cfg["preview_size"] = tuple(cfg.get("preview_size", (1280, 720))) @@ -556,10 +556,6 @@ def normalizar_config_runtime_seg(cfg: dict) -> dict: if "veto_labels" in det and not isinstance(det["veto_labels"], set): det["veto_labels"] = set(det["veto_labels"]) - fuser = grid.get("fuser", {}) or {} - if "y_range_m" in fuser: - fuser["y_range_m"] = tuple(fuser["y_range_m"]) - cfg["grid"] = grid segmentacao = cfg.get("segmentacao", {}) or {} @@ -615,6 +611,8 @@ def load_seg_config(): # Cópia profunda simples dos blocos aninhados. # Evita compartilhar dict interno entre chamadas. + cfg["telemetry"] = dict(VISUAL_DEFAULT_CONFIG["telemetry"]) + cfg["model_runtime"] = dict(VISUAL_DEFAULT_CONFIG["model_runtime"]) cfg["segmentacao"] = dict(VISUAL_DEFAULT_CONFIG["segmentacao"]) cfg["grid"] = { "grid_shape": VISUAL_DEFAULT_CONFIG["grid"]["grid_shape"], diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/main.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/main.py index 4ad0d187b..4c1a799da 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/main.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/main.py @@ -20,8 +20,13 @@ def main(): # Importar e configurar OpenCV antes dos módulos do Visual Worker, # pois eles podem carregar OpenCV internamente. import cv2 - cv2.setNumThreads(2) + opencv_threads = max(1, int(os.environ.get("VISUAL_OPENCV_THREADS", "4"))) + cv2.setNumThreads(opencv_threads) cv2.ocl.setUseOpenCL(False) + print( + f"[visual][PERF] OpenCV threads requested={opencv_threads} " + f"effective={cv2.getNumThreads()} preview_cache=True" + ) from shared.enums import VisualWorkerCommandType, TipoFrameCamera from shared.utils import encode_image_base64 @@ -53,6 +58,7 @@ def main(): time.sleep(0.1) def redis_callback(dados): + acao = None try: acao = VisualWorkerCommandType(dados.get("cmd", 0)) @@ -60,14 +66,14 @@ def main(): mx_id = dados.get("params") if mx_id is not None: iniciar_camera_manager(mx_id) - elif acao == VisualWorkerCommandType.CalibrarProfundidade: - n_frames = dados.get("params", 50) - get_camera_manager().gerar_grid_ref(n_frames) elif acao == VisualWorkerCommandType.AtualizarSaudeCamera: get_camera_manager().atualizar_saude_camera() elif acao == VisualWorkerCommandType.GetCameraFrame: tipo = TipoFrameCamera(dados.get("params", TipoFrameCamera.Rgb.value)) - frame = get_camera_manager().get_selected_frame(tipo) + frame = get_camera_manager().get_selected_frame( + tipo, + aguardar_atualizacao=True, + ) if frame is not None: base64_img = encode_image_base64(frame) if base64_img is not None: @@ -82,14 +88,11 @@ def main(): pasta = dados.get("params", {}).get("caminho", "frames_salvos") tipos = dados.get("params", {}).get("tipos", []) get_camera_manager().salvar_frames(tipos, nome, pasta) - elif acao == VisualWorkerCommandType.EnviarImagemMock: - caminho = dados.get("params", "") - get_camera_manager().segmentacao_manager.use_mock = caminho != "" - get_camera_manager().segmentacao_manager.img_mock = caminho else: mostrar_log(f"⚠️ Comando desconhecido: {acao.name}") except Exception as e: - mostrar_log(f"Erro ao processar comando {acao.name}: {e}") + acao_nome = getattr(acao, "name", str(dados.get("cmd", "desconhecido"))) + mostrar_log(f"Erro ao processar comando {acao_nome}: {e}") def inicializar(): mostrar_log("🚀 Iniciando Visual Worker...") diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/segmentacao_semantica.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/segmentacao_semantica.py index a3067e2b4..ac1ed865f 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/segmentacao_semantica.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/segmentacao_semantica.py @@ -2,7 +2,6 @@ from __future__ import annotations import time from dataclasses import dataclass, field -from enum import IntEnum from collections import deque from typing import Any, Dict, List, Optional, Sequence, Tuple @@ -12,23 +11,19 @@ import numpy as np from shared.enums import StatusCarroMapa -class ClassesSegmentacao(IntEnum): - NAONAVEGAVEL = 0 - NAVEGAVEL = 1 - - @dataclass class SegmentacaoConfig: """ - Configuração v1 do analisador semântico do Visual Worker. + Configuração v2 do analisador semântico do Visual Worker. Esta classe controla apenas análise de máscara. Não controla ONNX, câmera, render, preview, stream nem Redis. """ - # IDs esperados na máscara de segmentação. - id_nao_navegavel: int = int(ClassesSegmentacao.NAONAVEGAVEL) - id_navegavel: int = int(ClassesSegmentacao.NAVEGAVEL) + # IDs são OBRIGATORIAMENTE injetados do contrato ONNX v2. + # -1 faz o componente falhar fechado se alguém tentar inicializá-lo isolado. + id_nao_navegavel: int = -1 + id_navegavel: int = -1 # Scanlines usadas para estimar centro, largura e ângulo do corredor. # Frações verticais do frame: 0.0 = topo, 1.0 = base. @@ -42,13 +37,27 @@ class SegmentacaoConfig: ema_alpha_lat: float = 0.25 ema_alpha_conf: float = 0.20 - # Histórico para estabilizar status do corredor. - status_window_s: float = 1.5 + # Histórico ponderado para estabilizar status do corredor. + status_window_s: float = 0.90 + status_decay_tau_s: float = 0.35 status_expected_fps: float = 10.0 - # Prioridade do status vindo da cabeça auxiliar ONNX. - model_conf_accept: float = 0.70 - model_conf_soft: float = 0.45 + # Cabeça de status ONNX v2: confiança + margem top1-top2. + model_conf_accept: float = 0.72 + model_margin_accept: float = 0.15 + model_conf_soft: float = 0.55 + model_margin_soft: float = 0.06 + + # Troca rápida para transições muito seguras. + model_conf_fast: float = 0.90 + model_margin_fast: float = 0.30 + fast_switch_min_frames: int = 2 + + # Pesos do voto temporal. + status_weight_model: float = 1.35 + status_weight_model_soft: float = 1.00 + status_weight_heuristic: float = 0.75 + status_weight_agreement_bonus: float = 0.40 # Grid leve usada para pontuar corredor e fallback de centro. corridor_grid_rows: int = 6 @@ -70,15 +79,16 @@ class SegmentacaoConfig: # Debug leve: inclui métricas extras no retorno, sem criar imagens. include_debug: bool = True + include_model_probs: bool = False debug_timing: bool = False class SegmentacaoManager: """ - Analisador de segmentação semântica v1. + Analisador de segmentação semântica v2. Responsabilidade única: - pred_ids + aux_result -> dados_visuais + pred_ids + status_result -> dados_visuais Não renderiza. Não colore máscara. @@ -94,6 +104,14 @@ class SegmentacaoManager: ): self.config = self._normalizar_config(config) + if int(self.config.id_navegavel) < 0 or int(self.config.id_nao_navegavel) < 0: + raise RuntimeError( + "IDs semânticos não foram injetados pelo contrato ONNX v2: " + f"nav={self.config.id_navegavel} non_nav={self.config.id_nao_navegavel}" + ) + if int(self.config.id_navegavel) == int(self.config.id_nao_navegavel): + raise RuntimeError("IDs navegável/não-navegável são iguais no contrato.") + # Mantidos como metadados úteis, mas não usados no caminho quente. self.color_map = color_map self.classes = classes @@ -105,6 +123,8 @@ class SegmentacaoManager: self._status_hist = deque(maxlen=max_len) self._status_final_hist = deque(maxlen=2) + self._strong_candidate: Optional[StatusCarroMapa] = None + self._strong_candidate_count: int = 0 self._ema_ang: Optional[float] = None self._ema_lat: Optional[float] = None @@ -118,22 +138,22 @@ class SegmentacaoManager: # API principal # --------------------------------------------------------------------- - def segmentar(self, predictions: np.ndarray, aux_result: Optional[Dict[str, Any]] = None): + def segmentar(self, predictions: np.ndarray, status_result: Optional[Dict[str, Any]] = None): """ Compatibilidade nominal para o CameraManager atual. - Para a v1, o método preferido é analisar(). + No contrato v2, o método preferido é analisar(). Retorna: resultado, erro """ try: - resultado = self.analisar(predictions, aux_result=aux_result) + resultado = self.analisar(predictions, status_result=status_result) return resultado, None except Exception as e: self._ultimo_erro = f"Erro na análise de segmentação: {e}" return None, self._ultimo_erro - def analisar(self, predictions: np.ndarray, aux_result: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: + def analisar(self, predictions: np.ndarray, status_result: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: t0 = time.perf_counter() pred = self._validar_predictions(predictions) @@ -161,7 +181,7 @@ class SegmentacaoManager: pred=pred, mask_nav=mask_nav, corridor_score=corridor_score, - aux_result=aux_result, + status_result=status_result, ) t_status = time.perf_counter() @@ -182,10 +202,26 @@ class SegmentacaoManager: "status_corredor_anterior": int(status_before.value), "status_corredor_anterior_nome": str(status_before.name), + # Telemetria compacta da segunda cabeça. Não depende do debug pesado. + "status_fonte": str(status_debug.get("origem_status", "desconhecida")), + "status_modelo_id": (status_debug.get("modelo") or {}).get("label_id"), + "status_modelo_nome": (status_debug.get("modelo") or {}).get("status_nome"), + "status_modelo_conf": (status_debug.get("modelo") or {}).get("conf"), + "status_modelo_margin": (status_debug.get("modelo") or {}).get("margin"), + "status_modelo_entropy": (status_debug.get("modelo") or {}).get("entropy"), + "status_modelo_second_nome": (status_debug.get("modelo") or {}).get("second_name"), + "status_modelo_second_conf": (status_debug.get("modelo") or {}).get("second_conf"), + "status_fast_switch": bool(status_debug.get("fast_switch", False)), + "centros_corredor": centros, "larguras_px": larguras_px, } + if self.config.include_model_probs: + probs = (status_debug.get("modelo") or {}).get("probs") + if probs is not None: + dados_visuais["status_modelo_probs"] = probs + if self.config.include_debug: dados_visuais["status_corredor_debug"] = status_debug dados_visuais["debug_corredor"] = { @@ -670,29 +706,100 @@ class SegmentacaoManager: pred: np.ndarray, mask_nav: np.ndarray, corridor_score: float, - aux_result: Optional[Dict[str, Any]], + status_result: Optional[Dict[str, Any]], ): - model_status, model_conf, model_debug = self._status_from_aux(aux_result) + model_status, model_conf, model_margin, model_debug = self._status_from_model( + status_result + ) - heur_status, heur_debug = self._status_heuristico(pred, mask_nav, corridor_score) + heur_status, heur_debug = self._status_heuristico( + pred, mask_nav, corridor_score + ) origem = "heuristica" status_now = heur_status + weight = float(self.config.status_weight_heuristic) - if model_status is not None and model_conf >= self.config.model_conf_accept: + model_fast = ( + model_status is not None + and model_conf >= self.config.model_conf_fast + and model_margin >= self.config.model_margin_fast + ) + + model_accept = ( + model_status is not None + and model_conf >= self.config.model_conf_accept + and model_margin >= self.config.model_margin_accept + ) + + model_soft = ( + model_status is not None + and model_conf >= self.config.model_conf_soft + and model_margin >= self.config.model_margin_soft + ) + + if model_accept: status_now = model_status - origem = "modelo" - elif model_status is not None and model_conf >= self.config.model_conf_soft: - # Modelo com média confiança só vence quando a heurística está indefinida. + origem = "modelo_fast" if model_fast else "modelo" + weight = float(self.config.status_weight_model) + # Confiança e margem entram suavemente no peso, sem explodir a janela. + weight *= 0.75 + 0.25 * float(np.clip(model_conf, 0.0, 1.0)) + weight *= 0.85 + 0.15 * float(np.clip(model_margin, 0.0, 1.0)) + + if heur_status == model_status: + weight += float(self.config.status_weight_agreement_bonus) + origem += "_agree" + + elif model_soft: if heur_status == StatusCarroMapa.Indefinido: status_now = model_status origem = "modelo_soft" + weight = float(self.config.status_weight_model_soft) + elif heur_status == model_status: + status_now = model_status + origem = "modelo_soft_agree" + weight = ( + float(self.config.status_weight_model_soft) + + float(self.config.status_weight_agreement_bonus) + ) else: status_now = heur_status - origem = "heuristica_com_modelo_soft" + origem = "heuristica_modelo_soft_discorda" + + now = self._now() + self._status_hist.append((status_now, now, float(weight), origem)) + + # Transições realmente fortes não precisam aguardar toda a janela temporal. + fast_switch = False + if model_fast and model_status is not None: + if self._strong_candidate == model_status: + self._strong_candidate_count += 1 + else: + self._strong_candidate = model_status + self._strong_candidate_count = 1 + + if self._strong_candidate_count >= max(1, int(self.config.fast_switch_min_frames)): + status_final = model_status + fast_switch = True + # Re-semeia a janela com o estado forte para não voltar no frame seguinte. + self._status_hist.clear() + self._status_hist.append(( + status_final, + now, + float(self.config.status_weight_model) + + float(self.config.status_weight_agreement_bonus), + "fast_switch_seed", + )) + else: + status_final, weighted_scores = self._weighted_status() + else: + self._strong_candidate = None + self._strong_candidate_count = 0 + status_final, weighted_scores = self._weighted_status() + + if fast_switch: + weighted_scores = {status_final.name: 1.0} - self._status_hist.append((status_now, self._now())) - status_final = self._majority_status() self._status_final_hist.append(status_final) status_before = ( @@ -705,34 +812,64 @@ class SegmentacaoManager: "origem_status": origem, "modelo": model_debug, "heuristica": heur_debug, + "status_now": status_now.name, + "status_final": status_final.name, + "candidate_weight": round(float(weight), 4), + "weighted_scores": weighted_scores, + "fast_switch": bool(fast_switch), + "fast_candidate": ( + self._strong_candidate.name + if self._strong_candidate is not None + else None + ), + "fast_candidate_count": int(self._strong_candidate_count), } return status_now, status_final, status_before, debug - def _status_from_aux(self, aux_result: Optional[Dict[str, Any]]): - if not aux_result or aux_result.get("type") != "label": - return None, 0.0, { + def _status_from_model(self, status_result: Optional[Dict[str, Any]]): + if not status_result: + return None, 0.0, 0.0, { "disponivel": False, - "motivo": "sem_aux_result", + "motivo": "sem_status_result", } try: - label_id = int(aux_result.get("label_id", StatusCarroMapa.Indefinido.value)) - conf = float(aux_result.get("label_conf", 0.0)) - status = StatusCarroMapa(label_id) + label_name = str(status_result.get("label_name", "")).strip() + conf = float(status_result.get("confidence", 0.0)) + margin = float(status_result.get("margin", 0.0)) + entropy = float(status_result.get("entropy", 1.0)) - return status, conf, { + if not label_name: + raise ValueError("label_name vazio") + + # Mapeamento oficial é por NOME. A ordem dos IDs do treino pode mudar + # sem embaralhar a enum do controle. + status = StatusCarroMapa[label_name] + + debug = { "disponivel": True, "status": int(status.value), "status_nome": str(status.name), - "label_name": aux_result.get("label_name"), + "label_id": int(status_result.get("label_id", -1)), + "label_name": label_name, "conf": round(conf, 4), + "margin": round(margin, 4), + "entropy": round(entropy, 4), + "second_id": int(status_result.get("second_id", -1)), + "second_name": status_result.get("second_name"), + "second_conf": round(float(status_result.get("second_confidence", 0.0)), 4), } - except Exception as e: - return None, 0.0, { + if "probs" in status_result: + debug["probs"] = list(status_result["probs"]) + + return status, conf, margin, debug + + except Exception as exc: + return None, 0.0, 0.0, { "disponivel": False, - "motivo": f"aux_invalido: {e}", + "motivo": f"status_modelo_invalido: {exc}", } def _status_heuristico( @@ -827,33 +964,39 @@ class SegmentacaoManager: def _now() -> float: return time.monotonic() - def _majority_status(self, janela_s: Optional[float] = None) -> StatusCarroMapa: - J = self.config.status_window_s if janela_s is None else float(janela_s) + def _weighted_status(self) -> Tuple[StatusCarroMapa, Dict[str, float]]: + window_s = max(0.05, float(self.config.status_window_s)) + tau_s = max(0.05, float(self.config.status_decay_tau_s)) now = self._now() - while self._status_hist and (now - self._status_hist[0][1] > J): + while self._status_hist and (now - self._status_hist[0][1] > window_s): self._status_hist.popleft() if not self._status_hist: - return StatusCarroMapa.Direcionando + return StatusCarroMapa.Direcionando, {} - counts: Dict[StatusCarroMapa, int] = {} + scores: Dict[StatusCarroMapa, float] = {} + latest_ts: Dict[StatusCarroMapa, float] = {} - for status, _t in self._status_hist: - counts[status] = counts.get(status, 0) + 1 - - top = max(counts.values()) - tied = [status for status, count in counts.items() if count == top] + for status, ts, base_weight, _source in self._status_hist: + age = max(0.0, now - float(ts)) + decay = float(np.exp(-age / tau_s)) + score = max(0.01, float(base_weight)) * decay + scores[status] = scores.get(status, 0.0) + score + latest_ts[status] = max(latest_ts.get(status, 0.0), float(ts)) + top_score = max(scores.values()) + tied = [s for s, score in scores.items() if abs(score - top_score) <= 1e-9] if len(tied) == 1: - return tied[0] + winner = tied[0] + else: + winner = max(tied, key=lambda st: latest_ts.get(st, 0.0)) - # Desempate: status mais recente. - for status, _t in reversed(self._status_hist): - if status in tied: - return status - - return StatusCarroMapa.Direcionando + debug_scores = { + status.name: round(float(score), 5) + for status, score in sorted(scores.items(), key=lambda item: item[1], reverse=True) + } + return winner, debug_scores # --------------------------------------------------------------------- # Suavização @@ -879,6 +1022,8 @@ class SegmentacaoManager: def reset_estado_temporal(self): self._status_hist.clear() self._status_final_hist.clear() + self._strong_candidate = None + self._strong_candidate_count = 0 self._ema_ang = None self._ema_lat = None diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/visual_debug_renderer.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/visual_debug_renderer.py index 66b0d0e56..e300257c5 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/visual_debug_renderer.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/visual_debug_renderer.py @@ -57,9 +57,12 @@ class VisualDebugRenderer: return self.build_detection_overlay(rgb_frame, detections) if frame_type == TipoFrameCamera.Debug: - overlay = self.build_overlay(rgb_frame, pred_ids, alpha=alpha) + # Debug prefere costmap. Só monta overlay como fallback, evitando + # construir duas imagens completas para descartar uma delas. grid = self.build_costmap_debug(rgb_frame, snapshot) - return grid if grid is not None else overlay + if grid is not None: + return grid + return self.build_overlay(rgb_frame, pred_ids, alpha=alpha) return None @@ -67,12 +70,12 @@ class VisualDebugRenderer: if rgb_frame is None: return None - img = rgb_frame.copy() + if rgb_frame.shape[1::-1] != self.preview_size: + # cv2.resize já cria um novo buffer; copiar antes seria tráfego de + # memória inútil no caminho de preview. + return cv2.resize(rgb_frame, self.preview_size, interpolation=cv2.INTER_AREA) - if img.shape[1::-1] != self.preview_size: - img = cv2.resize(img, self.preview_size, interpolation=cv2.INTER_AREA) - - return img + return rgb_frame.copy() def build_segmentation_preview(self, pred_ids): if pred_ids is None or self.color_map is None: diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/visual_grid_builder.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/visual_grid_builder.py index 04db412a4..6a41e0991 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/visual_grid_builder.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/processamento/visual_grid_builder.py @@ -6,8 +6,6 @@ from typing import Any, Dict, Optional, Sequence, Tuple import cv2 import numpy as np -from visual_worker.processamento.segmentacao_semantica import ClassesSegmentacao - @dataclass class GridGeometryConfig: @@ -41,6 +39,10 @@ class DetectionGridConfig: @dataclass class GridConfidenceConfig: + # IDs vêm OBRIGATORIAMENTE do contrato ONNX e são injetados pelo CameraManager. + id_nao_navegavel: int = -1 + id_navegavel: int = -1 + valid_mm: Tuple[int, int] = (300, 10000) min_valid_frac: float = 0.30 conf_params: Tuple[float, float] = (0.30, 0.80) @@ -144,6 +146,14 @@ class VisualGridBuilder: def __init__(self, config: Optional[GridConfidenceConfig | Dict[str, Any]] = None): self.config = self._normalizar_config(config) + if int(self.config.id_navegavel) < 0 or int(self.config.id_nao_navegavel) < 0: + raise RuntimeError( + "VisualGridBuilder exige IDs do contrato ONNX v2: " + f"nav={self.config.id_navegavel} non_nav={self.config.id_nao_navegavel}" + ) + if int(self.config.id_navegavel) == int(self.config.id_nao_navegavel): + raise RuntimeError("Grid recebeu IDs navegável/não-navegável iguais.") + @staticmethod def _normalizar_config(config): if config is None: @@ -236,8 +246,8 @@ class VisualGridBuilder: if n == 0: continue - n_nav = np.count_nonzero(seg_block == ClassesSegmentacao.NAVEGAVEL.value) - n_naonav = np.count_nonzero(seg_block == ClassesSegmentacao.NAONAVEGAVEL.value) + n_nav = np.count_nonzero(seg_block == cfg.id_navegavel) + n_naonav = np.count_nonzero(seg_block == cfg.id_nao_navegavel) pct_navegavel[j, i] = n_nav / n pct_nao_navegavel[j, i] = n_naonav / n diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/utils.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/utils.py index 6529e4865..457cae6b9 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/utils.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/utils.py @@ -1,144 +1,31 @@ import cv2 import numpy as np -import cupy as cp -import math -from scipy import stats -from shared.enums import T_Code + def gerar_heatmap(depth_frame, dist_max_mm=3000.0): mask_valid = depth_frame <= dist_max_mm depth_normalized = np.zeros(depth_frame.shape, dtype=np.uint8) - depth_normalized[mask_valid] = (255 - ((depth_frame[mask_valid] / dist_max_mm) * 255)).astype(np.uint8) + depth_normalized[mask_valid] = ( + 255 - ((depth_frame[mask_valid] / dist_max_mm) * 255) + ).astype(np.uint8) heatmap = cv2.applyColorMap(depth_normalized, cv2.COLORMAP_JET) heatmap[~mask_valid] = [0, 0, 0] return heatmap -def calcular_threshold_anomalias(velocidade_ms=None, threshold_base=0.15, fator_sensibilidade=100): - """ - Calcula o threshold dinamicamente baseado na velocidade do robô. - - velocidade_ms: Velocidade em metros por segundo (m/s). - """ - try: - if velocidade_ms is None: - from shared.contexto_global_redis import ContextoGlobalRedis - velocidade_ms = ContextoGlobalRedis.get_contexto().get("Gerais", {}).get("velocidade_ms", 0.0) - threshold = threshold_base + (velocidade_ms * fator_sensibilidade) - return round(threshold, 3) - except Exception as e: - print(f"❌ Erro ao calcular thresold de anomalias: {e}") - return threshold_base - -def obter_regiao_solo(depth_frame, percentual_altura): - depth_frame = cp.asarray(depth_frame) # 🔥 Garante que está na GPU - - altura, largura = depth_frame.shape - linhas_solo = int((percentual_altura / 100) * altura) - y_inicio = altura - linhas_solo - - depth_solo = depth_frame[y_inicio:altura, :] # 🔥 Slice direto na GPU - - return depth_solo, y_inicio, linhas_solo, altura, largura - -def calcular_referencia(frames, metodo="moda"): - stack = np.stack(frames) - - if metodo == "moda": - ref = stats.mode(stack, axis=0, keepdims=True)[0][0] - elif metodo == "media": - ref = np.mean(stack, axis=0) - elif metodo == "mediana": - ref = np.median(stack, axis=0) - else: - raise ValueError("Método inválido. Use 'moda', 'media' ou 'mediana'.") - - return ref - -def calcular_inclinacao_solo(perfil, fov_h): - """ - Calcula o ângulo de inclinação do solo em graus. - - :param perfil: Lista de profundidades dos setores [z1, z2, ..., zn] - :param largura_total_metros: Largura total em metros do campo de visão - :return: Ângulo de inclinação em graus - """ - if len(perfil) < 2: - return 0 # Não dá pra calcular - - distancia_media = np.mean(perfil) - largura = 2 * math.tan(fov_h / 2) * distancia_media - - delta_z = perfil[-1] - perfil[0] - delta_x = largura - - angulo_rad = math.atan2(delta_z, delta_x) - angulo_graus = math.degrees(angulo_rad) - - return round(angulo_graus, 3) - -def calcular_profundidades_setores(depth_solo, n_setores, largura, distancia_max_mm): - depth_solo = cp.asarray(depth_solo) # 🔥 Garante que está na GPU - - largura_setor = largura // n_setores - profundidades = [] - - for i in range(n_setores): - x_inicio = i * largura_setor - x_fim = (i + 1) * largura_setor if (i < n_setores - 1) else largura - - setor = depth_solo[:, x_inicio:x_fim] - - # 🔥 Filtra profundidade válida - setor_valido = setor[(setor > 0) & (setor < distancia_max_mm)] - - z = cp.median(setor_valido).item() if setor_valido.size > 0 else 0 # 🔥 .item() p/ float - profundidades.append(z) - - return cp.asnumpy(cp.array(profundidades)), largura_setor - -def histogram2d_gpu(x, y, bins, range): - try: - if x.size == 0 or y.size == 0: - return cp.zeros(bins, dtype=cp.int32) - - x_bins, y_bins = bins - (x_min, x_max), (y_min, y_max) = range - - x_idx = cp.floor((x - x_min) / (x_max - x_min) * x_bins).astype(cp.int32) - y_idx = cp.floor((y - y_min) / (y_max - y_min) * y_bins).astype(cp.int32) - - mask = (x_idx >= 0) & (x_idx < x_bins) & (y_idx >= 0) & (y_idx < y_bins) - x_idx = x_idx[mask] - y_idx = y_idx[mask] - - if x_idx.size == 0: - return cp.zeros((x_bins, y_bins), dtype=cp.int32) - - linear_idx = x_idx * y_bins + y_idx - - if linear_idx.size == 0: - return cp.zeros((x_bins, y_bins), dtype=cp.int32) - - hist_flat = cp.bincount(linear_idx, minlength=x_bins * y_bins) - hist = hist_flat.reshape((x_bins, y_bins)) - return hist - except Exception as e: - print(f"Erro ao gerar histograma 2d: {e}") - return cp.zeros((x_bins, y_bins), dtype=cp.int32) def converter_valores_numpy(obj): if isinstance(obj, dict): return {k: converter_valores_numpy(v) for k, v in obj.items()} - elif isinstance(obj, list): + if isinstance(obj, list): return [converter_valores_numpy(v) for v in obj] - elif isinstance(obj, tuple): + if isinstance(obj, tuple): return tuple(converter_valores_numpy(v) for v in obj) - elif isinstance(obj, (np.integer, np.int32, np.int64)): + if isinstance(obj, (np.integer, np.int32, np.int64)): return int(obj) - elif isinstance(obj, (np.floating, np.float32, np.float64)): + if isinstance(obj, (np.floating, np.float32, np.float64)): return float(obj) - elif isinstance(obj, (np.bool_)): + if isinstance(obj, np.bool_): return bool(obj) - elif isinstance(obj, np.ndarray): + if isinstance(obj, np.ndarray): return obj.tolist() - else: - return obj + return obj diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/camera_manager.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/camera_manager.py index 83de8c489..c233ea788 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/camera_manager.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/camera_manager.py @@ -28,7 +28,7 @@ from visual_worker.utils import converter_valores_numpy from shared.perf_monitor import VisualPerfMonitor -CAMERA_MANAGER_VERSION = "production_v1_2026_09_12_runtime_priority_v3" +CAMERA_MANAGER_VERSION = "production_v1_2026_09_15_debug_geometry_sync_v4" PRODUCT_ASSEMBLY_SCHEMA = "multispec_module_params_assembly_v1" PRODUCT_MODULE_SCHEMA = "multispec_module_params_v3" @@ -2044,6 +2044,135 @@ class CameraManager: return True + @staticmethod + def _clamp01_debug(value, default=0.0): + try: + return max(0.0, min(1.0, float(value))) + except Exception: + return max(0.0, min(1.0, float(default))) + + def _resolver_geometria_debug_bicos(self, config, analise_debug=None): + """ + Resolve a geometria que o Debug deve desenhar. + + Autoridade, nesta ordem: + 1) estatísticas produzidas pelo WeedDetector na última decisão real; + 2) o próprio helper do WeedDetector, quando disponível; + 3) fallback local com o MESMO contrato: 0=topo, 1=pé. + + Assim o preview não mantém uma segunda implementação da lógica de atuação. + """ + cfg = config or {} + analise_debug = analise_debug or {} + + dados_visuais = analise_debug.get("dados_visuais", {}) or {} + estatisticas = dados_visuais.get("estatisticas", {}) or {} + estat_bicos = estatisticas.get("bicos", {}) or {} + estat_atuacao = estat_bicos.get("atuacao", {}) or {} + + ordem_invertida = bool( + estat_bicos.get( + "ordem_bicos_invertida", + cfg.get("inverter_ordem_bicos", False), + ) + ) + sentido_vertical_invertido = bool( + estat_bicos.get( + "sentido_vertical_invertido", + estat_atuacao.get( + "sentido_vertical_invertido", + cfg.get("inverter_sentido_vertical_bicos", False), + ), + ) + ) + + zona_ini = estat_atuacao.get("zona_atuacao_inicio_frac") + zona_fim = estat_atuacao.get("zona_atuacao_fim_frac") + eval_ini = estat_atuacao.get("zona_eval_inicio_frac") + eval_fim = estat_atuacao.get("zona_eval_fim_frac") + + # Se ainda não existe uma decisão real (startup, reset etc.), pede ao + # detector a mesma função usada pelo controle. Evita duplicar semântica. + if zona_ini is None or zona_fim is None: + detector = self.weed_detector + helper = getattr(detector, "_calcular_zona_atuacao_frac_from_config", None) + if callable(helper): + try: + zona_ini, zona_fim = helper(cfg) + except Exception: + zona_ini = None + zona_fim = None + + # Último fallback, mantido fail-soft para o Debug. Esta é a convenção + # oficial atual do WeedDetector: faixa é medida de cima para baixo. + if zona_ini is None or zona_fim is None: + faixa = self._clamp01_debug(cfg.get("faixa_atuacao_bicos", 0.70), 0.70) + area = self._clamp01_debug(cfg.get("area_atuacao_bicos", 0.10), 0.10) + area = max(0.001, area) + + zona_ini = faixa + zona_fim = min(1.0, faixa + area) + + if sentido_vertical_invertido: + zona_ini, zona_fim = 1.0 - zona_fim, 1.0 - zona_ini + + zona_ini = self._clamp01_debug(zona_ini, 0.0) + zona_fim = self._clamp01_debug(zona_fim, zona_ini) + if zona_fim < zona_ini: + zona_ini, zona_fim = zona_fim, zona_ini + + # Sem compensação de latência disponível, a zona avaliada coincide com + # a zona física. Quando o detector informar outra zona, mostramos ambas. + if eval_ini is None or eval_fim is None: + eval_ini, eval_fim = zona_ini, zona_fim + + eval_ini = self._clamp01_debug(eval_ini, zona_ini) + eval_fim = self._clamp01_debug(eval_fim, zona_fim) + if eval_fim < eval_ini: + eval_ini, eval_fim = eval_fim, eval_ini + + pulverizacao = dados_visuais.get("pulverizacao", {}) or {} + + return { + "zona_inicio_frac": float(zona_ini), + "zona_fim_frac": float(zona_fim), + "zona_eval_inicio_frac": float(eval_ini), + "zona_eval_fim_frac": float(eval_fim), + "ordem_bicos_invertida": bool(ordem_invertida), + "sentido_vertical_invertido": bool(sentido_vertical_invertido), + "pulverizacao": pulverizacao, + "fonte_detector": bool(estat_atuacao), + } + + @staticmethod + def _frac_vertical_para_pixels(inicio_frac, fim_frac, altura): + if altura <= 1: + return 0, 0 + + inicio = max(0.0, min(1.0, float(inicio_frac))) + fim = max(0.0, min(1.0, float(fim_frac))) + if fim < inicio: + inicio, fim = fim, inicio + + y0 = int(round(inicio * (altura - 1))) + y1 = int(round(fim * (altura - 1))) + y0 = max(0, min(altura - 1, y0)) + y1 = max(0, min(altura - 1, y1)) + return min(y0, y1), max(y0, y1) + + @staticmethod + def _intervalo_visual_bico(indice_bico, qtd_bicos, largura, ordem_invertida): + """Mapeia índice FÍSICO do bico para a coluna correspondente da imagem.""" + qtd = max(1, int(qtd_bicos)) + idx = int(indice_bico) + idx_visual = (qtd - 1 - idx) if ordem_invertida else idx + + x0 = int(round(idx_visual * largura / float(qtd))) + x1 = int(round((idx_visual + 1) * largura / float(qtd))) - 1 + x0 = max(0, min(largura - 1, x0)) + x1 = max(x0, min(largura - 1, x1)) + return x0, x1 + def get_debug_frame(self, mostrar=False, overlay_bgr=None): overlay = overlay_bgr if overlay_bgr is not None else self._ultimo_preview_overlay @@ -2057,12 +2186,21 @@ class CameraManager: "fps_loop": self._ultimo_loop_analise_fps, } + # A detecção publica controle final + análise sob _vida_lock. Tiramos o + # snapshot sob o mesmo lock para o Debug não misturar estados de ciclos + # diferentes enquanto a thread de detecção está atualizando os caches. + with self._vida_lock: + atuacao_bicos = dict(self._ultimo_controle or {}) + analise_debug = dict(self._ultima_analise or {}) + config_debug = dict(self.seg_config or {}) + # Não publica diretamente no cache: o preview producer faz a # publicação atômica depois de reduzir/comprimir o bundle inteiro. return self._montar_debug_overlay( overlay_bgr=overlay, - atuacao_bicos=self._ultimo_controle or {}, - config=self.seg_config or {}, + atuacao_bicos=atuacao_bicos, + config=config_debug, + analise_debug=analise_debug, metricas_perf=metricas_perf, mostrar=mostrar, ) @@ -2072,6 +2210,7 @@ class CameraManager: overlay_bgr, atuacao_bicos, config, + analise_debug=None, metricas_perf=None, mostrar=False, ): @@ -2080,6 +2219,10 @@ class CameraManager: return None metricas_perf = metricas_perf or {} + geometria = self._resolver_geometria_debug_bicos( + config=config, + analise_debug=analise_debug, + ) W, H = self._dbg_shape @@ -2104,25 +2247,68 @@ class CameraManager: self._dbg_img[...] = overlay_bgr qtd_bicos = int(config.get("qtd_bicos", self.qtd_bicos or 1) or 1) - zona_inicio = float(config.get("faixa_atuacao_bicos", 0.7)) - faixa_atuacao = float(config.get("area_atuacao_bicos", 0.1)) + qtd_bicos = max(1, qtd_bicos) - y_inicio = int((1.0 - zona_inicio) * H) - y_fim = int((1.0 - (zona_inicio + faixa_atuacao)) * H) - y0, y1 = min(y_inicio, y_fim), max(y_inicio, y_fim) + y0, y1 = self._frac_vertical_para_pixels( + geometria["zona_inicio_frac"], + geometria["zona_fim_frac"], + H, + ) + ey0, ey1 = self._frac_vertical_para_pixels( + geometria["zona_eval_inicio_frac"], + geometria["zona_eval_fim_frac"], + H, + ) + # ------------------------------------------------------------ + # Zona FÍSICA dos bicos. É a faixa configurada pelo operador. + # ------------------------------------------------------------ self._dbg_layer.fill(0) cv2.rectangle( self._dbg_layer, (0, y0), - (W, y1), + (W - 1, y1), (220, 220, 100), thickness=-1, ) - cv2.addWeighted(self._dbg_layer, 0.18, self._dbg_img, 0.82, 0, dst=self._dbg_img) - cv2.rectangle(self._dbg_img, (0, y0), (W, y1), (180, 180, 80), 2) + cv2.addWeighted( + self._dbg_layer, + 0.18, + self._dbg_img, + 0.82, + 0, + dst=self._dbg_img, + ) + cv2.rectangle( + self._dbg_img, + (0, y0), + (W - 1, y1), + (180, 180, 80), + 2, + ) + + # Zona REALMENTE avaliada pelo detector quando há antecipação por + # latência/velocidade. Se coincidir com a zona física, não duplica. + eval_diferente = abs(ey0 - y0) > 1 or abs(ey1 - y1) > 1 + if eval_diferente: + cv2.rectangle( + self._dbg_img, + (0, ey0), + (W - 1, ey1), + (255, 255, 0), + 2, + ) + cv2.putText( + self._dbg_img, + "EVAL", + (max(5, W - 60), max(16, min(H - 5, ey0 + 16))), + cv2.FONT_HERSHEY_SIMPLEX, + 0.45, + (255, 255, 0), + 1, + cv2.LINE_AA, + ) - largura_bico = W / float(max(qtd_bicos, 1)) cores = [ (0, 255, 0), (255, 0, 0), @@ -2134,33 +2320,64 @@ class CameraManager: (0, 0, 255), ] + ordem_invertida = bool(geometria["ordem_bicos_invertida"]) + + # ------------------------------------------------------------ + # Estado ON/OFF FINAL, já depois do gate de pulverização. + # Cada índice é o bico FÍSICO. A posição visual respeita + # inverter_ordem_bicos exatamente como o WeedDetector. + # ------------------------------------------------------------ self._dbg_layer.fill(0) - for i in range(qtd_bicos): - x0 = int(i * largura_bico) - x1 = int((i + 1) * largura_bico) + x0, x1 = self._intervalo_visual_bico( + i, + qtd_bicos, + W, + ordem_invertida, + ) cor = cores[i % len(cores)] - if atuacao_bicos.get(i, False): - cv2.rectangle(self._dbg_layer, (x0, y0), (x1, y1), cor, thickness=-1) + if bool(atuacao_bicos.get(i, False)): + cv2.rectangle( + self._dbg_layer, + (x0, y0), + (x1, y1), + cor, + thickness=-1, + ) - cv2.addWeighted(self._dbg_layer, 0.15, self._dbg_img, 0.85, 0, dst=self._dbg_img) + cv2.addWeighted( + self._dbg_layer, + 0.15, + self._dbg_img, + 0.85, + 0, + dst=self._dbg_img, + ) for i in range(qtd_bicos): - x0 = int(i * largura_bico) - x1 = int((i + 1) * largura_bico) + x0, x1 = self._intervalo_visual_bico( + i, + qtd_bicos, + W, + ordem_invertida, + ) cor = cores[i % len(cores)] - status = "ON" if atuacao_bicos.get(i, False) else "OFF" + status = "ON" if bool(atuacao_bicos.get(i, False)) else "OFF" cv2.rectangle(self._dbg_img, (x0, y0), (x1, y1), cor, 1) + + texto = f"Bico {i} {status}" + text_y = max(18, min(H - 6, y0 + 20)) cv2.putText( self._dbg_img, - f"Bico {i} {status}", - (x0 + 5, min(H - 5, y1 + 20)), + texto, + (x0 + 4, text_y), cv2.FONT_HERSHEY_SIMPLEX, - 0.55, + 0.48, cor, - 2, + 1, + cv2.LINE_AA, ) fps_infer = float(metricas_perf.get("fps_infer") or 0.0) @@ -2188,6 +2405,51 @@ class CameraManager: 2, ) + zona_txt = ( + f"Zona={geometria['zona_inicio_frac'] * 100:.0f}-" + f"{geometria['zona_fim_frac'] * 100:.0f}%" + ) + if eval_diferente: + zona_txt += ( + f" | Eval={geometria['zona_eval_inicio_frac'] * 100:.0f}-" + f"{geometria['zona_eval_fim_frac'] * 100:.0f}%" + ) + + zona_txt += ( + f" | V={'INV' if geometria['sentido_vertical_invertido'] else 'NORMAL'}" + f" H={'INV' if ordem_invertida else 'NORMAL'}" + ) + + cv2.putText( + self._dbg_img, + zona_txt, + (10, 88), + cv2.FONT_HERSHEY_SIMPLEX, + 0.48, + (230, 230, 230), + 1, + cv2.LINE_AA, + ) + + pulverizacao = geometria.get("pulverizacao", {}) or {} + if pulverizacao: + permitida = bool(pulverizacao.get("permitida", False)) + motivo = str(pulverizacao.get("motivo", "") or "") + gate_txt = f"Pulv: {'ON' if permitida else 'OFF'}" + if motivo and motivo != "ok": + gate_txt += f" ({motivo[:48]})" + + cv2.putText( + self._dbg_img, + gate_txt, + (10, 108), + cv2.FONT_HERSHEY_SIMPLEX, + 0.45, + (200, 255, 200) if permitida else (180, 180, 255), + 1, + cv2.LINE_AA, + ) + if mostrar: cv2.imshow("Debug Weed Worker", self._dbg_img) cv2.waitKey(1) diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/config.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/config.py index 3cbb17f1e..b76c82bc8 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/config.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/config.py @@ -256,7 +256,7 @@ WEED_DEFAULT_CONFIG = { # True no rover atual: # a aproximação física do alvo ocorre do pé para o topo da imagem. - "inverter_sentido_vertical_bicos": True, + "inverter_sentido_vertical_bicos": False, # Zona física de atuação no frame. # Ambos podem ser sobrescritos pela configuração da operação. diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/weed_detector.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/weed_detector.py index 166718d68..f5c48bbed 100644 --- a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/weed_detector.py +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/weed_worker/weed_detector.py @@ -25,11 +25,14 @@ class WeedDetector: - Desloca a memória pela distância real percorrida pelo rover. - Atua quando a evidência chega na zona física de atuação. - Convenção espacial: - - cell index 0 = região mais próxima da linha dos bicos. - - cell index maior = região mais distante, vista antes pela câmera. - - Conforme o rover anda, a memória se desloca de índices maiores - para índices menores. + Convenção espacial oficial: + - eixo vertical sempre usa coordenada de imagem: 0.0 = topo, 1.0 = pé. + - cell index 0 = topo da imagem. + - cell index N-1 = pé da imagem. + - inverter_sentido_vertical_bicos=False (rover atual): + alvo caminha topo -> pé; faixa=0.85, area=0.15 atua em 85%..100%. + - inverter_sentido_vertical_bicos=True: + alvo caminha pé -> topo; a mesma faixa é espelhada para 0%..15%. Este arquivo NÃO: - gera imagem; @@ -41,6 +44,7 @@ class WeedDetector: CONTRATO_OFICIAL = "target_binary" TARGET_ID = 1 + VERSION = "weed_detector_v2_2026_09_15_vertical_contract_fix" def __init__(self, config: Optional[dict] = None): # O CameraManager já mantém um snapshot de configuração. Quando ele é @@ -76,8 +80,14 @@ class WeedDetector: self.config.get("inverter_ordem_bicos", False) ) + # Convenção oficial: + # False = fluxo normal topo -> pé (rover atual). + # True = fluxo invertido pé -> topo. + # + # IMPORTANTE: esta flag altera tanto o sentido da memória espacial + # quanto a posição vertical da zona física de atuação. self.inverter_sentido_vertical_bicos = bool( - self.config.get("inverter_sentido_vertical_bicos", True) + self.config.get("inverter_sentido_vertical_bicos", False) ) self.detector_debug_perf = bool( @@ -168,11 +178,14 @@ class WeedDetector: 0.0 = topo da imagem 1.0 = pé da imagem - Normal: + Normal (inverter_sentido_vertical_bicos=False): faixa=0.90, area=0.10 -> 90% até 100%, embaixo. - Com inverter_sentido_vertical_bicos=True: + Invertido (inverter_sentido_vertical_bicos=True): faixa=0.90, area=0.10 -> 0% até 10%, em cima. + + Portanto, para o rover atual, em que o alvo se aproxima do topo para + o pé da imagem, a configuração correta é False. """ faixa = float(cfg.get("faixa_atuacao_bicos", 0.70) or 0.70) @@ -918,11 +931,12 @@ class WeedDetector: k_shift = max(0.0, k_shift) # A zona configurada representa a posição física fixa dos bicos. - # Com faixa=0,85, área=0,10 e inversão vertical, 85%..95% da imagem - # original vira 5%..15% na convenção interna. Essa zona NÃO muda de - # parado até a velocidade de referência. Somente a parcela de - # velocidade ACIMA da referência antecipa a avaliação para compensar - # o tempo de resposta do sistema. + # Exemplo no rover atual (sentido normal / False): + # faixa=0,85, área=0,15 -> 85%..100%, no pé da imagem. + # No sentido invertido / True a mesma faixa é espelhada para 0%..15%. + # Essa zona NÃO muda de parado até a velocidade de referência. Somente + # a parcela de velocidade ACIMA da referência antecipa a avaliação para + # compensar o tempo de resposta do sistema. vel_referencia = max( 0.0, float(cfg.get("velocidade_referencia_atuacao_mps", 0.65) or 0.65), @@ -1282,19 +1296,46 @@ class WeedDetector: vel_norm: float, cfg: dict, ): - # Convenção legada: - # zona_inicio e zona_altura são frações medidas a partir da parte inferior da imagem. - y_inicio = int((1.0 - zona_inicio) * h) - y_fim = int((1.0 - (zona_inicio + zona_altura)) * h) + """ + ROI do fallback legado usando a MESMA convenção do caminho oficial. - y_top = min(y_inicio, y_fim) - y_bot = max(y_inicio, y_fim) + Coordenada vertical: + 0.0 = topo + 1.0 = pé + False: alvo topo -> pé; 0.85 + 0.15 => faixa inferior. + True : alvo pé -> topo; a mesma faixa é espelhada para o topo. + + Os argumentos zona_inicio/zona_altura são mantidos na assinatura por + compatibilidade, mas o cálculo usa os valores de cfg para garantir que + exista uma única fonte de verdade. + """ + _ = zona_inicio, zona_altura + + zona_ini_frac, zona_fim_frac = self._calcular_zona_atuacao_frac_from_config(cfg) + + y_top = int(np.floor(zona_ini_frac * h)) + y_bot = int(np.ceil(zona_fim_frac * h)) + + y_top = max(0, min(h, y_top)) + y_bot = max(y_top, min(h, y_bot)) + + # Antecipação legada em pixels. A avaliação deve se deslocar para a + # região que o alvo ocupa ANTES de chegar aos bicos: + # fluxo topo -> pé : antecipa para cima (subtrai y) + # fluxo pé -> topo : antecipa para baixo (soma y) roi_shift_per_v = float(cfg.get("k_roi_shift_px_per_vnorm", 24.0)) - shift = int(roi_shift_per_v * self._clamp01(vel_norm)) + shift = int(max(0.0, roi_shift_per_v) * self._clamp01(vel_norm)) - y_top = max(0, y_top - shift) - y_bot = max(y_top, min(h, y_bot - shift)) + if getattr(self, "inverter_sentido_vertical_bicos", False): + y_top = min(h, y_top + shift) + y_bot = min(h, y_bot + shift) + else: + y_top = max(0, y_top - shift) + y_bot = max(0, y_bot - shift) + + if y_bot < y_top: + y_top, y_bot = y_bot, y_top return int(y_top), int(y_bot), int(shift) diff --git a/Python/OAK/datasets/oak-d/_0_capture_corridor.py b/Python/OAK/datasets/oak-d/_0_capture_corridor.py new file mode 100644 index 000000000..6cf9206c6 --- /dev/null +++ b/Python/OAK/datasets/oak-d/_0_capture_corridor.py @@ -0,0 +1,851 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +Capture OAK-D Lite - Agri Corridor Dataset +========================================== + +Captura RGB para o dataset da camera frontal. + +Contrato desta etapa: + - roda de dentro da pasta oak-d/ + - captura o mesmo dominio de video usado no runtime: + OAK-D Lite -> ColorCamera 1080p -> cam.video -> host + - salva SOMENTE a imagem RGB original em PNG, sem resize e sem normalizacao + - destino padrao: + dataset/brutas/ + - resize existe apenas para PREVIEW da interface e nunca toca a imagem salva + +Teclas: + SPACE / S : salvar frame atual + A : ligar/desligar auto-save + E : alternar exposicao AUTO/MANUAL + + / = : aumentar ISO no modo manual + - : diminuir ISO no modo manual + M : aumentar exposicao no modo manual + N : diminuir exposicao no modo manual + F : alternar foco AUTO/MANUAL + ] : aumentar foco manual + [ : diminuir foco manual + V : alternar fullscreen + Q / ESC : sair + +Exemplo: + python _0_capture.py + +Auto-save a cada 1 segundo: + python _0_capture.py --auto --interval 1.0 + +Outra pasta: + python _0_capture.py --out dataset/brutas +""" + +from __future__ import annotations + +import argparse +import time +from concurrent.futures import ThreadPoolExecutor +from datetime import datetime +from pathlib import Path +from typing import Optional, Tuple + +import cv2 +import depthai as dai +import numpy as np + + +CAPTURE_W = 1920 +CAPTURE_H = 1080 +CAPTURE_FPS = 30.0 + +DEFAULT_OUT = Path("dataset") / "brutas" + +DEFAULT_EXPOSURE_US = 6000 +DEFAULT_ISO = 400 +DEFAULT_FOCUS = 130 + +EXPOSURE_MIN_US = 100 +EXPOSURE_MAX_US = 30000 +EXPOSURE_STEP_US = 500 + +ISO_MIN = 100 +ISO_MAX = 1600 +ISO_STEP = 50 + +FOCUS_MIN = 0 +FOCUS_MAX = 255 +FOCUS_STEP = 5 + +WINDOW_NAME = "OAK-D Lite | Dataset Corredor Agricola" +CAPTURE_SESSION_ID = datetime.now().strftime("%Y%m%d_%H%M%S") + + +def timestamp_name() -> str: + """ + Nome novo: + img_sYYYYMMDD_HHMMSS__YYYYMMDD_HHMMSS_micro.png + + O primeiro timestamp identifica a SESSAO de captura (uma execução do script). + O segundo identifica o frame. + + Isso permite ao split manter uma mesma sessão inteira em train OU val. + """ + frame_ts = datetime.now().strftime("%Y%m%d_%H%M%S_%f") + return f"img_s{CAPTURE_SESSION_ID}__{frame_ts}" + + +def save_png(path: Path, frame_bgr: np.ndarray) -> Tuple[bool, str]: + path.parent.mkdir(parents=True, exist_ok=True) + + ok = cv2.imwrite( + str(path), + frame_bgr, + [cv2.IMWRITE_PNG_COMPRESSION, 3], + ) + + return bool(ok), str(path) + + +def fit_inside(image: np.ndarray, max_w: int, max_h: int) -> np.ndarray: + h, w = image.shape[:2] + + if w <= 0 or h <= 0: + return image + + scale = min( + float(max_w) / float(w), + float(max_h) / float(h), + ) + scale = max(scale, 1e-6) + + new_w = max(1, int(round(w * scale))) + new_h = max(1, int(round(h * scale))) + + interp = cv2.INTER_AREA if scale < 1.0 else cv2.INTER_LINEAR + return cv2.resize(image, (new_w, new_h), interpolation=interp) + + +def live_quality(frame_bgr: np.ndarray) -> dict: + """ + QA apenas visual. + Nunca rejeita nem altera a imagem salva. + """ + small = fit_inside(frame_bgr, 640, 360) + gray = cv2.cvtColor(small, cv2.COLOR_BGR2GRAY) + + mean = float(gray.mean()) + dark_pct = float((gray <= 5).mean() * 100.0) + bright_pct = float((gray >= 250).mean() * 100.0) + sharpness = float(cv2.Laplacian(gray, cv2.CV_64F).var()) + + warnings = [] + + if mean < 35: + warnings.append("MUITO ESCURA") + elif mean > 225: + warnings.append("MUITO CLARA") + + if dark_pct > 20.0: + warnings.append("CLIP PRETO") + + if bright_pct > 8.0: + warnings.append("CLIP BRANCO") + + if sharpness < 35.0: + warnings.append("FOCO/BLUR?") + + return { + "mean": mean, + "dark_pct": dark_pct, + "bright_pct": bright_pct, + "sharpness": sharpness, + "warnings": warnings, + } + + +def put_line( + canvas: np.ndarray, + text: str, + xy: Tuple[int, int], + scale: float = 0.65, + thickness: int = 1, + color=(235, 235, 235), +): + cv2.putText( + canvas, + text, + xy, + cv2.FONT_HERSHEY_SIMPLEX, + scale, + color, + thickness, + cv2.LINE_AA, + ) + + +class CameraControls: + def __init__( + self, + queue, + exposure_us: int, + iso: int, + focus: int, + ): + self.queue = queue + + self.exposure_auto = True + self.focus_auto = True + + self.exposure_us = int(exposure_us) + self.iso = int(iso) + self.focus = int(focus) + + def send_initial_auto(self): + ctrl = dai.CameraControl() + + try: + ctrl.setAutoExposureEnable() + except Exception: + pass + + try: + ctrl.setAutoWhiteBalanceLock(False) + except Exception: + pass + + try: + ctrl.setAutoFocusMode( + dai.CameraControl.AutoFocusMode.CONTINUOUS_VIDEO + ) + ctrl.setAutoFocusTrigger() + except Exception: + pass + + self.queue.send(ctrl) + + def set_exposure_auto(self): + ctrl = dai.CameraControl() + + try: + ctrl.setAutoExposureEnable() + except Exception as exc: + print(f"[WARN] Nao consegui ativar AE: {exc}") + return + + self.queue.send(ctrl) + self.exposure_auto = True + print("[CAM] Exposicao: AUTO") + + def set_exposure_manual(self): + self.exposure_us = int( + np.clip( + self.exposure_us, + EXPOSURE_MIN_US, + EXPOSURE_MAX_US, + ) + ) + self.iso = int( + np.clip( + self.iso, + ISO_MIN, + ISO_MAX, + ) + ) + + ctrl = dai.CameraControl() + ctrl.setManualExposure( + self.exposure_us, + self.iso, + ) + self.queue.send(ctrl) + + self.exposure_auto = False + + print( + f"[CAM] Exposicao: MANUAL | " + f"{self.exposure_us} us | ISO {self.iso}" + ) + + def toggle_exposure(self): + if self.exposure_auto: + self.set_exposure_manual() + else: + self.set_exposure_auto() + + def adjust_exposure(self, delta_us: int): + if self.exposure_auto: + print("[CAM] Ajuste de exposicao ignorado: pressione E para modo MANUAL.") + return + + self.exposure_us += int(delta_us) + self.set_exposure_manual() + + def adjust_iso(self, delta_iso: int): + if self.exposure_auto: + print("[CAM] Ajuste de ISO ignorado: pressione E para modo MANUAL.") + return + + self.iso += int(delta_iso) + self.set_exposure_manual() + + def set_focus_auto(self): + ctrl = dai.CameraControl() + + try: + ctrl.setAutoFocusMode( + dai.CameraControl.AutoFocusMode.CONTINUOUS_VIDEO + ) + ctrl.setAutoFocusTrigger() + except Exception as exc: + print(f"[WARN] Autofocus nao disponivel: {exc}") + return + + self.queue.send(ctrl) + self.focus_auto = True + print("[CAM] Foco: AUTO CONTINUOUS") + + def set_focus_manual(self): + self.focus = int( + np.clip( + self.focus, + FOCUS_MIN, + FOCUS_MAX, + ) + ) + + ctrl = dai.CameraControl() + + try: + ctrl.setManualFocus(self.focus) + except Exception as exc: + print(f"[WARN] Foco manual nao disponivel: {exc}") + return + + self.queue.send(ctrl) + self.focus_auto = False + print(f"[CAM] Foco: MANUAL | lens={self.focus}") + + def toggle_focus(self): + if self.focus_auto: + self.set_focus_manual() + else: + self.set_focus_auto() + + def adjust_focus(self, delta: int): + if self.focus_auto: + print("[CAM] Ajuste de foco ignorado: pressione F para modo MANUAL.") + return + + self.focus += int(delta) + self.set_focus_manual() + + +def create_pipeline(fps: float) -> dai.Pipeline: + pipeline = dai.Pipeline() + + cam_rgb = pipeline.createColorCamera() + cam_rgb.setBoardSocket(dai.CameraBoardSocket.CAM_A) + cam_rgb.setResolution( + dai.ColorCameraProperties.SensorResolution.THE_1080_P + ) + cam_rgb.setFps(float(fps)) + cam_rgb.setInterleaved(False) + + # Mesmo contrato do runtime atual. + cam_rgb.setColorOrder( + dai.ColorCameraProperties.ColorOrder.RGB + ) + + xout = pipeline.createXLinkOut() + xout.setStreamName("rgb") + + # IMPORTANTE: VIDEO, nao preview. + cam_rgb.video.link(xout.input) + + control_in = pipeline.createXLinkIn() + control_in.setStreamName("control") + control_in.out.link(cam_rgb.inputControl) + + return pipeline + + +def build_canvas( + frame_bgr: np.ndarray, + last_saved: Optional[np.ndarray], + *, + controls: CameraControls, + auto_save: bool, + interval_s: float, + saved_count: int, + out_dir: Path, + fps_view: float, + quality: dict, +) -> np.ndarray: + canvas_w = 1600 + canvas_h = 900 + + canvas = np.zeros( + (canvas_h, canvas_w, 3), + dtype=np.uint8, + ) + + live = fit_inside(frame_bgr, 1150, 780) + lh, lw = live.shape[:2] + + live_x = 20 + live_y = 60 + + canvas[ + live_y:live_y + lh, + live_x:live_x + lw, + ] = live + + panel_x = 1200 + + put_line( + canvas, + "OAK-D LITE / CORREDOR", + (panel_x, 60), + scale=0.78, + thickness=2, + ) + + put_line( + canvas, + "Fonte: VIDEO 1920x1080", + (panel_x, 100), + ) + + put_line( + canvas, + f"View FPS: {fps_view:.1f}", + (panel_x, 130), + ) + + put_line( + canvas, + f"Salvas: {saved_count}", + (panel_x, 160), + ) + + exp_text = ( + "AUTO" + if controls.exposure_auto + else f"MANUAL {controls.exposure_us}us ISO{controls.iso}" + ) + + focus_text = ( + "AUTO" + if controls.focus_auto + else f"MANUAL {controls.focus}" + ) + + put_line( + canvas, + f"Exposure: {exp_text}", + (panel_x, 205), + ) + + put_line( + canvas, + f"Focus: {focus_text}", + (panel_x, 235), + ) + + auto_text = ( + f"ON ({interval_s:.2f}s)" + if auto_save + else "OFF" + ) + + put_line( + canvas, + f"Auto-save: {auto_text}", + (panel_x, 265), + ) + + put_line( + canvas, + f"Luma mean: {quality['mean']:.1f}", + (panel_x, 315), + ) + put_line( + canvas, + f"Dark clip: {quality['dark_pct']:.1f}%", + (panel_x, 345), + ) + put_line( + canvas, + f"White clip: {quality['bright_pct']:.1f}%", + (panel_x, 375), + ) + put_line( + canvas, + f"Sharpness: {quality['sharpness']:.1f}", + (panel_x, 405), + ) + + if quality["warnings"]: + put_line( + canvas, + "QA: " + " | ".join(quality["warnings"]), + (panel_x, 440), + scale=0.55, + thickness=2, + color=(0, 180, 255), + ) + else: + put_line( + canvas, + "QA: OK", + (panel_x, 440), + scale=0.62, + thickness=2, + color=(80, 230, 80), + ) + + put_line(canvas, "SPACE/S salvar", (panel_x, 510)) + put_line(canvas, "A auto-save", (panel_x, 540)) + put_line(canvas, "E auto/manual exp", (panel_x, 570)) + put_line(canvas, "+/- ISO manual", (panel_x, 600)) + put_line(canvas, "M/N exposicao manual", (panel_x, 630)) + put_line(canvas, "F auto/manual foco", (panel_x, 660)) + put_line(canvas, "[/] foco manual", (panel_x, 690)) + put_line(canvas, "V fullscreen", (panel_x, 720)) + put_line(canvas, "Q/ESC sair", (panel_x, 750)) + + if last_saved is not None: + thumb = fit_inside(last_saved, 360, 100) + th, tw = thumb.shape[:2] + + tx = panel_x + ty = 780 + + if ty + th <= canvas_h and tx + tw <= canvas_w: + canvas[ + ty:ty + th, + tx:tx + tw, + ] = thumb + + put_line( + canvas, + f"Saida: {out_dir}", + (20, 30), + scale=0.62, + color=(200, 220, 255), + ) + + return canvas + + +def main(): + parser = argparse.ArgumentParser( + description=( + "Captura PNG lossless da OAK-D Lite para dataset " + "de corredor agricola." + ) + ) + + parser.add_argument( + "--out", + default=str(DEFAULT_OUT), + help="Fila de imagens brutas (somente PNGs).", + ) + + parser.add_argument( + "--fps", + type=float, + default=CAPTURE_FPS, + help="FPS da camera.", + ) + + parser.add_argument( + "--auto", + action="store_true", + help="Inicia com auto-save ligado.", + ) + + parser.add_argument( + "--interval", + type=float, + default=1.0, + help="Intervalo do auto-save em segundos.", + ) + + parser.add_argument( + "--exposure_us", + type=int, + default=DEFAULT_EXPOSURE_US, + help="Exposicao inicial quando entrar em modo manual.", + ) + + parser.add_argument( + "--iso", + type=int, + default=DEFAULT_ISO, + help="ISO inicial quando entrar em modo manual.", + ) + + parser.add_argument( + "--focus", + type=int, + default=DEFAULT_FOCUS, + help="Lens position inicial quando entrar em foco manual.", + ) + + args = parser.parse_args() + + out_dir = Path(args.out) + out_dir.mkdir(parents=True, exist_ok=True) + + saved_count = len(list(out_dir.glob("*.png"))) + + print("=" * 72) + print("OAK-D Lite | Captura Dataset Corredor Agricola") + print("=" * 72) + print(f"Saida : {out_dir.resolve()}") + print("Formato : PNG lossless") + print(f"Resolucao : {CAPTURE_W}x{CAPTURE_H}") + print("Fonte : ColorCamera.video") + print(f"FPS camera : {args.fps}") + print(f"Ja existem : {saved_count} PNGs") + print(f"Sessao : {CAPTURE_SESSION_ID}") + print("=" * 72) + print("SPACE/S salvar | A auto-save | E exposure | F foco | Q sair") + + pipeline = create_pipeline(args.fps) + + saver = ThreadPoolExecutor( + max_workers=1, + thread_name_prefix="png_saver", + ) + + pending = [] + last_saved: Optional[np.ndarray] = None + + auto_save = bool(args.auto) + interval_s = max(0.10, float(args.interval)) + last_auto_save_t = 0.0 + + fps_view = 0.0 + fps_frames = 0 + fps_t0 = time.perf_counter() + + fullscreen = False + + try: + with dai.Device(pipeline) as device: + control_queue = device.getInputQueue("control") + + rgb_queue = device.getOutputQueue( + "rgb", + maxSize=2, + blocking=False, + ) + + controls = CameraControls( + control_queue, + exposure_us=args.exposure_us, + iso=args.iso, + focus=args.focus, + ) + + controls.send_initial_auto() + + cv2.namedWindow( + WINDOW_NAME, + cv2.WINDOW_NORMAL, + ) + cv2.resizeWindow( + WINDOW_NAME, + 1600, + 900, + ) + + while True: + packet = rgb_queue.get() + frame = packet.getCvFrame() + + if ( + frame is None + or frame.ndim != 3 + or frame.shape[2] != 3 + ): + print( + f"[WARN] Frame invalido: " + f"{getattr(frame, 'shape', None)}" + ) + continue + + h, w = frame.shape[:2] + if (w, h) != (CAPTURE_W, CAPTURE_H): + print( + f"[WARN] Frame veio {w}x{h}, " + f"esperado {CAPTURE_W}x{CAPTURE_H}. " + "Nao sera redimensionado silenciosamente." + ) + + now_perf = time.perf_counter() + + fps_frames += 1 + dt_fps = now_perf - fps_t0 + if dt_fps >= 1.0: + fps_view = fps_frames / dt_fps + fps_frames = 0 + fps_t0 = now_perf + + quality = live_quality(frame) + + def request_save(current_frame: np.ndarray): + nonlocal saved_count, last_saved + + name = timestamp_name() + ".png" + path = out_dir / name + + snapshot = current_frame.copy() + + future = saver.submit( + save_png, + path, + snapshot, + ) + pending.append((future, path)) + + last_saved = snapshot + saved_count += 1 + + print( + f"[CAPTURE] #{saved_count} -> {path}" + ) + + if auto_save: + if ( + last_auto_save_t <= 0.0 + or now_perf - last_auto_save_t >= interval_s + ): + request_save(frame) + last_auto_save_t = now_perf + + still_pending = [] + for future, path in pending: + if not future.done(): + still_pending.append((future, path)) + continue + + try: + ok, saved_path = future.result() + except Exception as exc: + print( + f"[ERRO] Falha salvando {path}: {exc}" + ) + continue + + if not ok: + print( + f"[ERRO] cv2.imwrite retornou False: " + f"{saved_path}" + ) + + pending = still_pending + + canvas = build_canvas( + frame, + last_saved, + controls=controls, + auto_save=auto_save, + interval_s=interval_s, + saved_count=saved_count, + out_dir=out_dir, + fps_view=fps_view, + quality=quality, + ) + + cv2.imshow( + WINDOW_NAME, + canvas, + ) + + key = cv2.waitKey(1) & 0xFF + + if key in (ord("q"), 27): + print("[INFO] Encerrando captura...") + break + + if key in (ord("s"), ord(" ")): + request_save(frame) + + elif key == ord("a"): + auto_save = not auto_save + last_auto_save_t = 0.0 + print( + f"[CAPTURE] Auto-save " + f"{'ON' if auto_save else 'OFF'} " + f"| interval={interval_s:.2f}s" + ) + + elif key == ord("e"): + controls.toggle_exposure() + + elif key in (ord("+"), ord("=")): + controls.adjust_iso(ISO_STEP) + + elif key == ord("-"): + controls.adjust_iso(-ISO_STEP) + + elif key == ord("m"): + controls.adjust_exposure(EXPOSURE_STEP_US) + + elif key == ord("n"): + controls.adjust_exposure(-EXPOSURE_STEP_US) + + elif key == ord("f"): + controls.toggle_focus() + + elif key == ord("]"): + controls.adjust_focus(FOCUS_STEP) + + elif key == ord("["): + controls.adjust_focus(-FOCUS_STEP) + + elif key == ord("v"): + fullscreen = not fullscreen + cv2.setWindowProperty( + WINDOW_NAME, + cv2.WND_PROP_FULLSCREEN, + ( + cv2.WINDOW_FULLSCREEN + if fullscreen + else cv2.WINDOW_NORMAL + ), + ) + + finally: + cv2.destroyAllWindows() + + print( + f"[INFO] Aguardando {len(pending)} PNG(s) " + "que ja foram solicitados..." + ) + + saver.shutdown(wait=True) + + failures = 0 + for future, path in pending: + try: + ok, _ = future.result() + if not ok: + failures += 1 + except Exception as exc: + failures += 1 + print(f"[ERRO] {path}: {exc}") + + print("=" * 72) + print("Captura encerrada.") + print(f"Total contabilizado : {saved_count}") + print(f"Falhas finais : {failures}") + print(f"Saida : {out_dir.resolve()}") + print("=" * 72) + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/oak-d/_1_annotate_corridor.py b/Python/OAK/datasets/oak-d/_1_annotate_corridor.py new file mode 100644 index 000000000..8fc03a1a8 --- /dev/null +++ b/Python/OAK/datasets/oak-d/_1_annotate_corridor.py @@ -0,0 +1,2376 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +_1_annotate_corridor.py +======================= + +Editor definitivo da etapa BRUTAS para a camera frontal OAK-D Lite. + +Objetivo por imagem: + 1) criar/editar mascara binaria do corredor: + - nao navegavel + - navegavel + 2) definir o estado global do corredor: + - classes vindas de config.json -> label_classes + 3) derivar automaticamente o grupo pela mascara final: + - naonavegavel + - navegavel + - naonavegavel_navegavel + 4) arquivar a amostra pronta em dataset/original/group//: + images/.png + masks/.png # RGB colorida conforme labelmap + labels/.json + +A mascara nova nasce 100% NAO NAVEGAVEL. +O operador normalmente desenha apenas o corredor NAVEGAVEL. + +Ferramentas: + P = poligono + clique esquerdo -> adiciona ponto + clique direito -> remove ultimo ponto + ENTER -> fecha/preenche + + B = pincel + arrastar botao esquerdo -> pinta classe selecionada + arrastar botao direito -> pinta classe oposta + + F = balde / flood fill + clique esquerdo -> preenche regiao conectada com classe selecionada + + C = troca classe da mascara + [ / ] = diminui/aumenta pincel + Z = undo + Y = redo + R = reset para 100% nao navegavel + O = alterna visualizacao do overlay + +Estado global: + 1..9 / 0 = seleciona label conforme config["label_classes"] + A label NAO salva automaticamente. + +Fluxo: + S ou SPACE = SALVAR mascara + label e ir para a proxima + A / seta esquerda = anterior + D / seta direita = proxima sem salvar + Q / ESC = sair + +Execucao: + Rode de dentro da pasta oak-d: + + python _1_annotate_corridor.py + +Estrutura esperada: + oak-d/ + config.json + dataset/ + labelmap.txt + brutas/ # fila: somente PNGs ainda nao mapeadas + img_....png + original/ + group/ + navegavel/ + naonavegavel/ + naonavegavel_navegavel/ + +Observacoes: + - mascara BRUTA e salva em PNG RGB colorida, conforme labelmap; + - a imagem original nunca e recomprimida; ela e copiada/movida byte a byte; + - modo padrao MOVE: somente remove a PNG de brutas depois que image+mask+label finais existem; + - --mode copy preserva a PNG em brutas, mas amostras ja arquivadas sao puladas por padrao; + - esta etapa nao escolhe train/val. Ela apenas organiza o dataset original por grupo semantico. +""" + +from __future__ import annotations + +import argparse +import json +import os +import re +import shutil +import time +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from typing import Dict, List, Optional, Sequence, Tuple + +import cv2 +import numpy as np + + +# ============================================================================= +# Defaults +# ============================================================================= + +WINDOW_NAME = "Agrobot | Editor Corredor OAK-D Lite" + +IMG_EXTS = { + ".png", + ".jpg", + ".jpeg", + ".bmp", + ".webp", +} + +DEFAULT_BRUTAS_DIR = Path("dataset") / "brutas" +DEFAULT_ORIGINAL_GROUP_ROOT = Path("dataset") / "original" / "group" +DEFAULT_LABELMAP = Path("dataset") / "labelmap.txt" +DEFAULT_CONFIG = Path("config.json") + +MAX_WINDOW_W = 1800 +MAX_WINDOW_H = 1000 +PANEL_W = 430 +HEADER_H = 72 +FOOTER_H = 92 + +DEFAULT_OVERLAY_ALPHA = 0.42 +DEFAULT_BRUSH_RADIUS = 24 + +MASK_SCHEMA = "agrobot_corridor_mask_rgb_v1" +LABEL_SCHEMA = "agrobot_corridor_state_v2" + +# Cores de fallback apenas para visualizacao. +FALLBACK_COLORS_RGB = { + 0: (210, 55, 55), + 1: (45, 215, 70), +} + + +# ============================================================================= +# Estruturas +# ============================================================================= + +@dataclass +class SampleItem: + image_path: Path + mask_path: Path + label_path: Path + base: str + + +@dataclass +class ViewTransform: + x0: int = 0 + y0: int = 0 + width: int = 1 + height: int = 1 + scale: float = 1.0 + + +# ============================================================================= +# Utilitarios +# ============================================================================= + +def natural_key(text: str): + parts = re.split(r"(\d+)", str(text)) + return [ + int(p) if p.isdigit() else p.lower() + for p in parts + ] + + +def now_iso() -> str: + return datetime.now().isoformat(timespec="seconds") + + +def atomic_write_json(path: Path, data: dict) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_suffix(path.suffix + ".tmp") + + with tmp.open("w", encoding="utf-8") as f: + json.dump( + data, + f, + ensure_ascii=False, + indent=2, + ) + + os.replace(tmp, path) + + +def atomic_write_png_color(path: Path, mask_bgr: np.ndarray) -> None: + """Salva a máscara BRUTA colorida, lossless, conforme as cores exatas do labelmap.""" + path.parent.mkdir(parents=True, exist_ok=True) + + if mask_bgr.dtype != np.uint8: + mask_bgr = mask_bgr.astype(np.uint8) + + if mask_bgr.ndim != 3 or mask_bgr.shape[2] != 3: + raise ValueError( + f"Máscara colorida deve ser HxWx3 uint8; recebido {mask_bgr.shape}" + ) + + tmp = path.with_name(path.stem + ".__tmp__.png") + + ok = cv2.imwrite( + str(tmp), + mask_bgr, + [cv2.IMWRITE_PNG_COMPRESSION, 3], + ) + + if not ok: + try: + if tmp.exists(): + tmp.unlink() + except Exception: + pass + raise IOError(f"cv2.imwrite retornou False para {tmp}") + + os.replace(tmp, path) + + +def safe_relative(path: Path, root: Path) -> str: + try: + return str( + path.resolve().relative_to(root.resolve()) + ).replace("\\", "/") + except Exception: + return str(path).replace("\\", "/") + + +def read_json_safe(path: Path) -> Optional[dict]: + if not path.is_file(): + return None + + try: + with path.open("r", encoding="utf-8") as f: + value = json.load(f) + return value if isinstance(value, dict) else None + except Exception: + return None + + +def list_images(folder: Path) -> List[Path]: + if not folder.is_dir(): + raise FileNotFoundError(f"Pasta de imagens nao encontrada: {folder}") + + files = [ + p + for p in folder.iterdir() + if p.is_file() + and p.suffix.lower() in IMG_EXTS + ] + + files.sort(key=lambda p: natural_key(p.name)) + return files + + +# ============================================================================= +# Labelmap +# ============================================================================= + +def load_labelmap_classes( + path: Path, +) -> Dict[int, Tuple[str, Tuple[int, int, int]]]: + """ + Retorna: + id -> (nome, cor RGB) + + Formatos aceitos: + naonavegavel: 220,40,40 + navegavel: 40,220,80 + + 0 naonavegavel 220 40 40 + 1 navegavel 40 220 80 + + 0:naonavegavel + 1:navegavel + """ + if not path.is_file(): + raise FileNotFoundError( + f"labelmap obrigatorio nao encontrado: {path}" + ) + + classes: Dict[int, Tuple[str, Tuple[int, int, int]]] = {} + next_id = 0 + + with path.open("r", encoding="utf-8") as f: + for raw_line in f: + s = raw_line.strip() + + if not s or s.startswith("#"): + continue + + cid: Optional[int] = None + name: Optional[str] = None + color: Optional[Tuple[int, int, int]] = None + + left = s.split(":", 1)[0].strip() + + # Formato nome: R,G,B :: ... + if ":" in s and not left.isdigit(): + name_part, rest = s.split(":", 1) + name = name_part.strip() + + color_txt = ( + rest.split("::", 1)[0] + .strip() + .strip(":") + ) + + parts = [ + p.strip() + for p in color_txt.split(",") + if p.strip() + ] + + if len(parts) >= 3: + try: + color = tuple( + int(float(x)) + for x in parts[:3] + ) + except Exception: + color = None + + else: + parts = ( + s.replace(",", " ") + .replace(":", " ") + .split() + ) + + if len(parts) >= 2 and parts[0].isdigit(): + cid = int(parts[0]) + name = parts[1] + + if len(parts) >= 5: + try: + color = tuple( + int(float(x)) + for x in parts[2:5] + ) + except Exception: + color = None + + elif parts: + name = parts[0] + + if not name: + continue + + lname = name.strip().lower() + if lname in { + "ignore", + "void", + "background_ignore", + }: + continue + + if cid is None: + cid = next_id + + next_id = max(next_id, cid + 1) + + if color is None: + color = FALLBACK_COLORS_RGB.get( + int(cid), + (255, 255, 255), + ) + + classes[int(cid)] = ( + str(name), + tuple(map(int, color)), + ) + + if len(classes) != 2: + raise RuntimeError( + "Este editor e especializado em segmentacao binaria " + f"navegavel/nao navegavel. labelmap={classes}" + ) + + return classes + + +def class_color_bgr( + classes: Dict[int, Tuple[str, Tuple[int, int, int]]], + cid: int, +) -> Tuple[int, int, int]: + _name, rgb = classes[int(cid)] + return ( + int(rgb[2]), + int(rgb[1]), + int(rgb[0]), + ) + + +def resolve_nav_classes( + classes: Dict[int, Tuple[str, Tuple[int, int, int]]], + main_class_name: str, +) -> Tuple[int, int]: + target = str(main_class_name).strip().lower() + + exact = [ + cid + for cid, (name, _rgb) in classes.items() + if str(name).strip().lower() == target + ] + + if len(exact) != 1: + # Compatibilidade historica do projeto. + aliases = { + "navegavel", + "navegável", + "nav", + "navigable", + } + + exact = [ + cid + for cid, (name, _rgb) in classes.items() + if str(name).strip().lower() in aliases + ] + + if len(exact) != 1: + raise RuntimeError( + f"Nao consegui resolver classe navegavel. " + f"main_class_name={main_class_name!r} classes={classes}" + ) + + nav_id = int(exact[0]) + non_nav = [ + int(cid) + for cid in classes + if int(cid) != nav_id + ] + + if len(non_nav) != 1: + raise RuntimeError( + f"Nao consegui resolver classe nao navegavel: {classes}" + ) + + return nav_id, non_nav[0] + + +# ============================================================================= +# Mask IO +# ============================================================================= + +def colored_mask_to_ids( + mask_bgr: np.ndarray, + classes: Dict[int, Tuple[str, Tuple[int, int, int]]], +) -> np.ndarray: + h, w = mask_bgr.shape[:2] + ids = np.full( + (h, w), + 255, + dtype=np.uint8, + ) + + for cid in classes: + bgr = np.array( + class_color_bgr(classes, cid), + dtype=np.uint8, + ) + match = np.all( + mask_bgr == bgr[None, None, :], + axis=2, + ) + ids[match] = int(cid) + + if np.any(ids == 255): + unknown = int((ids == 255).sum()) + raise ValueError( + f"Mascara colorida possui {unknown} pixels " + "que nao casam exatamente com o labelmap." + ) + + return ids + + +def load_mask_ids( + path: Path, + expected_hw: Tuple[int, int], + classes: Dict[int, Tuple[str, Tuple[int, int, int]]], +) -> np.ndarray: + raw = cv2.imread( + str(path), + cv2.IMREAD_UNCHANGED, + ) + + if raw is None: + raise FileNotFoundError(f"Nao consegui abrir mask: {path}") + + if raw.ndim == 2: + ids = raw.astype( + np.uint8, + copy=False, + ) + elif raw.ndim == 3 and raw.shape[2] >= 3: + ids = colored_mask_to_ids( + raw[:, :, :3], + classes, + ) + else: + raise ValueError( + f"Formato de mask invalido: shape={raw.shape}" + ) + + if ids.shape != expected_hw: + raise ValueError( + f"Mask {path.name} shape={ids.shape} " + f"mas imagem shape={expected_hw}. " + "Nao farei resize silencioso." + ) + + valid_ids = set(int(x) for x in np.unique(ids)) + allowed = set(int(cid) for cid in classes) + + if not valid_ids.issubset(allowed): + raise ValueError( + f"Mask {path.name} possui IDs invalidos " + f"{sorted(valid_ids - allowed)}; permitidos={sorted(allowed)}" + ) + + return ids.copy() + + +def ids_to_color( + ids: np.ndarray, + classes: Dict[int, Tuple[str, Tuple[int, int, int]]], +) -> np.ndarray: + out = np.zeros( + (ids.shape[0], ids.shape[1], 3), + dtype=np.uint8, + ) + + for cid in classes: + out[ids == int(cid)] = class_color_bgr( + classes, + int(cid), + ) + + return out + + +# ============================================================================= +# Editor +# ============================================================================= + +class CorridorAnnotationEditor: + TOOL_POLYGON = "POLIGONO" + TOOL_BRUSH = "PINCEL" + TOOL_FILL = "BALDE" + + OVERLAY_NAV_ONLY = 0 + OVERLAY_BOTH = 1 + OVERLAY_OFF = 2 + + def __init__( + self, + *, + dataset_root: Path, + samples: Sequence[SampleItem], + states: Sequence[str], + classes: Dict[int, Tuple[str, Tuple[int, int, int]]], + nav_id: int, + non_nav_id: int, + original_group_root: Path, + archive_mode: str, + archived_before: int, + overlay_alpha: float, + brush_radius: int, + start_index: Optional[int], + ): + self.dataset_root = dataset_root + self.samples = list(samples) + self.states = [str(x) for x in states] + self.classes = classes + + self.nav_id = int(nav_id) + self.non_nav_id = int(non_nav_id) + self.original_group_root = Path(original_group_root) + self.archive_mode = str(archive_mode).lower() + if self.archive_mode not in {"move", "copy"}: + raise ValueError(f"archive_mode invalido: {self.archive_mode}") + self.archived_before = int(archived_before) + self.saved_this_session = 0 + self.finished = False + self.initial_queue_size = len(self.samples) + + self.resume = False + self.overlay_alpha = float( + np.clip( + overlay_alpha, + 0.0, + 1.0, + ) + ) + self.brush_radius = max( + 2, + int(brush_radius), + ) + + self.idx = 0 + + self.image: Optional[np.ndarray] = None + self.mask: Optional[np.ndarray] = None + self.original_loaded_mask: Optional[np.ndarray] = None + + self.selected_state: Optional[int] = None + self.loaded_state: Optional[int] = None + + self.selected_class = self.nav_id + self.tool = self.TOOL_POLYGON + self.overlay_mode = self.OVERLAY_NAV_ONLY + + self.polygon_points: List[Tuple[int, int]] = [] + + self.undo_stack: List[np.ndarray] = [] + self.redo_stack: List[np.ndarray] = [] + + self.mask_dirty = False + self.label_dirty = False + + self.view = ViewTransform() + self.last_canvas_shape = ( + MAX_WINDOW_H, + MAX_WINDOW_W, + ) + + self.mouse_down = False + self.mouse_button: Optional[int] = None + self.brush_action_started = False + self.mouse_src = (0, 0) + + self.message = "" + self.message_until = 0.0 + + self._completed_cache: Dict[str, bool] = {} + + if start_index is not None: + self.idx = max( + 0, + min( + len(self.samples) - 1, + int(start_index), + ), + ) + + cv2.namedWindow( + WINDOW_NAME, + cv2.WINDOW_NORMAL, + ) + cv2.resizeWindow( + WINDOW_NAME, + MAX_WINDOW_W, + MAX_WINDOW_H, + ) + cv2.setMouseCallback( + WINDOW_NAME, + self.on_mouse, + ) + + self.load_current() + + # ------------------------------------------------------------------------- + # Completion / navigation + # ------------------------------------------------------------------------- + + def sample_complete( + self, + item: SampleItem, + refresh: bool = False, + ) -> bool: + if not refresh and item.base in self._completed_cache: + return self._completed_cache[item.base] + + ok = False + + if ( + item.mask_path.is_file() + and item.label_path.is_file() + ): + try: + img = cv2.imread( + str(item.image_path), + cv2.IMREAD_COLOR, + ) + + if img is not None: + _ = load_mask_ids( + item.mask_path, + img.shape[:2], + self.classes, + ) + + lab = read_json_safe( + item.label_path, + ) + + lid = ( + int(lab["label_id"]) + if lab is not None + and "label_id" in lab + else -1 + ) + + ok = ( + 0 <= lid < len(self.states) + ) + except Exception: + ok = False + + self._completed_cache[item.base] = ok + return ok + + def find_first_incomplete(self) -> Optional[int]: + for i, item in enumerate(self.samples): + if not self.sample_complete(item): + return i + return None + + def completed_count(self) -> int: + return self.archived_before + self.saved_this_session + + def can_discard_current(self) -> bool: + return not ( + self.mask_dirty + or self.label_dirty + or self.polygon_points + ) + + def navigate( + self, + delta: int, + ): + delta = -1 if int(delta) < 0 else 1 + + if ( + self.mask_dirty + or self.label_dirty + or self.polygon_points + ): + now = time.monotonic() + pending = getattr( + self, + "_pending_nav", + None, + ) + + same_request = ( + isinstance(pending, tuple) + and len(pending) == 2 + and int(pending[0]) == delta + and now - float(pending[1]) <= 2.0 + ) + + if not same_request: + self._pending_nav = ( + delta, + now, + ) + direction = ( + "anterior" + if delta < 0 + else "proxima" + ) + self.flash( + f"Alteracoes nao salvas. " + f"Repita o comando para ir para {direction} " + "DESCARTANDO, ou S/SPACE para salvar.", + seconds=2.2, + ) + return + + self._pending_nav = None + self.flash( + "Alteracoes descartadas.", + seconds=1.0, + ) + + else: + self._pending_nav = None + + self.idx = max( + 0, + min( + len(self.samples) - 1, + self.idx + delta, + ), + ) + + self.load_current() + + # ------------------------------------------------------------------------- + # Load / save + # ------------------------------------------------------------------------- + + def load_current(self): + if not self.samples: + raise RuntimeError("Nenhuma imagem em dataset/brutas") + + item = self.samples[self.idx] + + image = cv2.imread( + str(item.image_path), + cv2.IMREAD_COLOR, + ) + + if image is None: + raise FileNotFoundError( + f"Nao consegui abrir imagem: {item.image_path}" + ) + + self.image = image + + if item.mask_path.is_file(): + self.mask = load_mask_ids( + item.mask_path, + image.shape[:2], + self.classes, + ) + else: + self.mask = np.full( + image.shape[:2], + self.non_nav_id, + dtype=np.uint8, + ) + + self.original_loaded_mask = self.mask.copy() + + existing = read_json_safe( + item.label_path, + ) + + self.selected_state = None + self.loaded_state = None + + if existing is not None: + try: + lid = int(existing.get("label_id")) + except Exception: + lid = -1 + + if 0 <= lid < len(self.states): + self.selected_state = lid + self.loaded_state = lid + + self.selected_class = self.nav_id + self.tool = self.TOOL_POLYGON + + self.polygon_points = [] + + self.undo_stack = [] + self.redo_stack = [] + + self.mask_dirty = False + self.label_dirty = False + + self.mouse_down = False + self.mouse_button = None + self.brush_action_started = False + + status = ( + "existente carregado" + if self.sample_complete(item, refresh=True) + else "novo/incompleto" + ) + + self.flash( + f"{item.image_path.name} | {status}", + seconds=1.4, + ) + + def derive_group(self) -> str: + assert self.mask is not None + + present = set(int(x) for x in np.unique(self.mask)) + allowed = {self.nav_id, self.non_nav_id} + if not present.issubset(allowed): + raise ValueError( + f"Mascara possui IDs inesperados: {sorted(present - allowed)}" + ) + + if present == {self.non_nav_id}: + return "naonavegavel" + if present == {self.nav_id}: + return "navegavel" + if present == {self.nav_id, self.non_nav_id}: + return "naonavegavel_navegavel" + + raise ValueError(f"Nao consegui derivar grupo da mascara: IDs={sorted(present)}") + + def final_paths(self, item: SampleItem, group: str): + root = self.original_group_root / group + image_dir = root / "images" + mask_dir = root / "masks" + label_dir = root / "labels" + for folder in (image_dir, mask_dir, label_dir): + folder.mkdir(parents=True, exist_ok=True) + + return ( + image_dir / item.image_path.name, + mask_dir / f"{item.base}.png", + label_dir / f"{item.base}.json", + ) + + def build_label_record( + self, + item: SampleItem, + group: str, + final_image: Path, + final_mask: Path, + final_label: Path, + ) -> dict: + assert self.selected_state is not None + assert self.mask is not None + assert self.image is not None + + nav_pct = float( + (self.mask == self.nav_id).mean() + * 100.0 + ) + + return { + "schema": LABEL_SCHEMA, + "estado_corredor": self.states[self.selected_state], + "label_id": int(self.selected_state), + "states": list(self.states), + + "base": item.base, + "group": group, + "image": safe_relative( + final_image, + self.dataset_root, + ), + "mask": safe_relative( + final_mask, + self.dataset_root, + ), + "label": safe_relative( + final_label, + self.dataset_root, + ), + "source_bruta": safe_relative( + item.image_path, + self.dataset_root, + ), + + "mask_schema": MASK_SCHEMA, + "mask_format": "rgb_color_png", + "mask_dtype": "uint8", + "mask_shape_hw": [ + int(self.mask.shape[0]), + int(self.mask.shape[1]), + ], + "mask_classes": { + str(cid): { + "name": self.classes[cid][0], + "rgb": list(self.classes[cid][1]), + } + for cid in sorted(self.classes) + }, + "nav_class_id": int(self.nav_id), + "non_nav_class_id": int(self.non_nav_id), + "navegavel_pct": nav_pct, + + "annotated_at": now_iso(), + "source": "agri_corridor_editor", + } + + def save_and_next(self): + item = self.samples[self.idx] + + if self.image is None or self.mask is None: + return + + if self.polygon_points: + self.flash( + "Ha poligono aberto. ENTER confirma ou BACKSPACE cancela." + ) + return + + if self.selected_state is None: + self.flash( + "Selecione o estado global com 1..9/0 antes de salvar." + ) + return + + valid_ids = set(int(x) for x in np.unique(self.mask)) + allowed = set(int(cid) for cid in self.classes) + + if not valid_ids.issubset(allowed): + self.flash( + f"ERRO mask IDs invalidos: {sorted(valid_ids - allowed)}", + seconds=3.0, + ) + return + + nav_pct = float((self.mask == self.nav_id).mean() * 100.0) + + try: + group = self.derive_group() + final_image, final_mask, final_label = self.final_paths(item, group) + + # Nunca sobrescreve silenciosamente uma amostra já arquivada. + collisions = [ + str(p) for p in (final_image, final_mask, final_label) if p.exists() + ] + if collisions: + raise FileExistsError( + "Amostra ja existe no destino final: " + ", ".join(collisions) + ) + + mask_bgr = ids_to_color(self.mask, self.classes) + record = self.build_label_record( + item, + group, + final_image, + final_mask, + final_label, + ) + + created = [] + try: + # Copia a imagem sem recomprimir. O source só é removido no fim. + final_image.parent.mkdir(parents=True, exist_ok=True) + shutil.copy2(item.image_path, final_image) + created.append(final_image) + + atomic_write_png_color(final_mask, mask_bgr) + created.append(final_mask) + + atomic_write_json(final_label, record) + created.append(final_label) + + if not all(p.is_file() for p in (final_image, final_mask, final_label)): + raise IOError("Commit incompleto: image/mask/label final nao confirmado") + + if self.archive_mode == "move": + item.image_path.unlink() + + except Exception: + # Como colisões são bloqueadas antes, tudo em created pertence a este commit. + for path in reversed(created): + try: + if path.exists(): + path.unlink() + except Exception: + pass + raise + + except Exception as exc: + self.flash( + f"ERRO salvando: {exc}", + seconds=4.0, + ) + return + + print( + f"[SAVE] {item.base} | group={group} | " + f"label={self.states[self.selected_state]} | " + f"nav={nav_pct:.2f}% | mode={self.archive_mode}" + ) + + self.saved_this_session += 1 + self.mask_dirty = False + self.label_dirty = False + + # A fila é consumida durante a sessão em MOVE e também em COPY. + # No COPY o arquivo físico permanece em brutas, mas na próxima execução + # collect_samples pula bases já arquivadas por padrão. + self.samples.pop(self.idx) + + if not self.samples: + self.finished = True + self.flash( + "Fila de brutas concluida. Todas as amostras desta sessao foram arquivadas.", + seconds=4.0, + ) + return + + self.idx = min(self.idx, len(self.samples) - 1) + self.load_current() + + # ------------------------------------------------------------------------- + # History + # ------------------------------------------------------------------------- + + def push_undo(self): + assert self.mask is not None + + self.undo_stack.append( + self.mask.copy() + ) + + if len(self.undo_stack) > 40: + self.undo_stack.pop(0) + + self.redo_stack.clear() + + def undo(self): + if not self.undo_stack: + self.flash("Nada para desfazer.") + return + + assert self.mask is not None + + self.redo_stack.append( + self.mask.copy() + ) + self.mask = self.undo_stack.pop() + self.mask_dirty = True + self.polygon_points = [] + + def redo(self): + if not self.redo_stack: + self.flash("Nada para refazer.") + return + + assert self.mask is not None + + self.undo_stack.append( + self.mask.copy() + ) + self.mask = self.redo_stack.pop() + self.mask_dirty = True + self.polygon_points = [] + + def reset_non_nav(self): + assert self.mask is not None + + self.push_undo() + self.mask.fill( + self.non_nav_id + ) + self.mask_dirty = True + self.polygon_points = [] + + self.flash( + "Mascara resetada para 100% nao navegavel." + ) + + # ------------------------------------------------------------------------- + # Tools + # ------------------------------------------------------------------------- + + def opposite_class( + self, + cid: int, + ) -> int: + return ( + self.non_nav_id + if int(cid) == self.nav_id + else self.nav_id + ) + + def toggle_class(self): + self.selected_class = self.opposite_class( + self.selected_class + ) + self.flash( + f"Classe de pintura: {self.classes[self.selected_class][0]}" + ) + + def set_tool( + self, + tool: str, + ): + self.tool = tool + self.polygon_points = [] + self.mouse_down = False + self.brush_action_started = False + self.flash(f"Ferramenta: {tool}") + + def commit_polygon(self): + if len(self.polygon_points) < 3: + self.flash( + "Poligono precisa de pelo menos 3 pontos." + ) + return + + assert self.mask is not None + + self.push_undo() + + poly = np.array( + self.polygon_points, + dtype=np.int32, + ) + + cv2.fillPoly( + self.mask, + [poly], + int(self.selected_class), + ) + + self.polygon_points = [] + self.mask_dirty = True + + def cancel_polygon(self): + self.polygon_points = [] + self.flash( + "Poligono aberto cancelado." + ) + + def paint_brush( + self, + src_x: int, + src_y: int, + cid: int, + ): + assert self.mask is not None + + cv2.circle( + self.mask, + (int(src_x), int(src_y)), + int(self.brush_radius), + int(cid), + -1, + cv2.LINE_8, + ) + + self.mask_dirty = True + + def flood_fill( + self, + src_x: int, + src_y: int, + ): + assert self.mask is not None + + y = int(src_y) + x = int(src_x) + + old = int(self.mask[y, x]) + new = int(self.selected_class) + + if old == new: + self.flash( + "Balde: regiao ja possui essa classe." + ) + return + + self.push_undo() + + # floodFill trabalha diretamente em uint8. + tmp = self.mask.copy() + + cv2.floodFill( + tmp, + None, + (x, y), + newVal=new, + loDiff=0, + upDiff=0, + flags=4, + ) + + self.mask = tmp + self.mask_dirty = True + + # ------------------------------------------------------------------------- + # Mouse mapping + # ------------------------------------------------------------------------- + + def window_to_source( + self, + x: int, + y: int, + ) -> Optional[Tuple[int, int]]: + if self.image is None: + return None + + v = self.view + + if ( + x < v.x0 + or y < v.y0 + or x >= v.x0 + v.width + or y >= v.y0 + v.height + ): + return None + + sx = int( + round( + (x - v.x0) + / max(v.scale, 1e-8) + ) + ) + sy = int( + round( + (y - v.y0) + / max(v.scale, 1e-8) + ) + ) + + h, w = self.image.shape[:2] + + sx = int( + np.clip( + sx, + 0, + w - 1, + ) + ) + sy = int( + np.clip( + sy, + 0, + h - 1, + ) + ) + + return sx, sy + + def source_to_window( + self, + x: int, + y: int, + ) -> Tuple[int, int]: + v = self.view + + return ( + int(round(v.x0 + x * v.scale)), + int(round(v.y0 + y * v.scale)), + ) + + def on_mouse( + self, + event, + x, + y, + flags, + param, + ): + src = self.window_to_source( + int(x), + int(y), + ) + + if src is None: + if event in { + cv2.EVENT_LBUTTONUP, + cv2.EVENT_RBUTTONUP, + }: + self.mouse_down = False + self.mouse_button = None + self.brush_action_started = False + return + + sx, sy = src + self.mouse_src = ( + sx, + sy, + ) + + # ------------------------------------------------------------- + # Polygon + # ------------------------------------------------------------- + + if self.tool == self.TOOL_POLYGON: + if event == cv2.EVENT_LBUTTONDOWN: + self.polygon_points.append( + (sx, sy) + ) + + elif event == cv2.EVENT_RBUTTONDOWN: + if self.polygon_points: + self.polygon_points.pop() + + return + + # ------------------------------------------------------------- + # Fill + # ------------------------------------------------------------- + + if self.tool == self.TOOL_FILL: + if event == cv2.EVENT_LBUTTONDOWN: + self.flood_fill( + sx, + sy, + ) + return + + # ------------------------------------------------------------- + # Brush + # ------------------------------------------------------------- + + if self.tool == self.TOOL_BRUSH: + if event in { + cv2.EVENT_LBUTTONDOWN, + cv2.EVENT_RBUTTONDOWN, + }: + self.mouse_down = True + self.mouse_button = event + self.brush_action_started = True + + self.push_undo() + + cid = ( + self.selected_class + if event == cv2.EVENT_LBUTTONDOWN + else self.opposite_class( + self.selected_class + ) + ) + + self.paint_brush( + sx, + sy, + cid, + ) + + elif ( + event == cv2.EVENT_MOUSEMOVE + and self.mouse_down + ): + cid = ( + self.selected_class + if self.mouse_button + == cv2.EVENT_LBUTTONDOWN + else self.opposite_class( + self.selected_class + ) + ) + + self.paint_brush( + sx, + sy, + cid, + ) + + elif event in { + cv2.EVENT_LBUTTONUP, + cv2.EVENT_RBUTTONUP, + }: + self.mouse_down = False + self.mouse_button = None + self.brush_action_started = False + + # ------------------------------------------------------------------------- + # Render + # ------------------------------------------------------------------------- + + def overlay_image(self) -> np.ndarray: + assert self.image is not None + assert self.mask is not None + + if self.overlay_mode == self.OVERLAY_OFF: + return self.image.copy() + + mask_color = ids_to_color( + self.mask, + self.classes, + ) + + blended = cv2.addWeighted( + self.image, + 1.0 - self.overlay_alpha, + mask_color, + self.overlay_alpha, + 0.0, + ) + + if self.overlay_mode == self.OVERLAY_BOTH: + return blended + + # Default: mostra apenas navegavel. + out = self.image.copy() + nav_region = ( + self.mask == self.nav_id + ) + out[nav_region] = blended[nav_region] + return out + + def build_canvas(self) -> np.ndarray: + assert self.image is not None + assert self.mask is not None + + canvas = np.zeros( + ( + MAX_WINDOW_H, + MAX_WINDOW_W, + 3, + ), + dtype=np.uint8, + ) + + left_w = MAX_WINDOW_W - PANEL_W + image_area_h = ( + MAX_WINDOW_H + - HEADER_H + - FOOTER_H + ) + + src_h, src_w = self.image.shape[:2] + + scale = min( + left_w / float(src_w), + image_area_h / float(src_h), + ) + + disp_w = max( + 1, + int(round(src_w * scale)), + ) + disp_h = max( + 1, + int(round(src_h * scale)), + ) + + x0 = max( + 0, + (left_w - disp_w) // 2, + ) + y0 = HEADER_H + max( + 0, + (image_area_h - disp_h) // 2, + ) + + self.view = ViewTransform( + x0=x0, + y0=y0, + width=disp_w, + height=disp_h, + scale=scale, + ) + + overlay = self.overlay_image() + + display = cv2.resize( + overlay, + (disp_w, disp_h), + interpolation=cv2.INTER_AREA, + ) + + canvas[ + y0:y0 + disp_h, + x0:x0 + disp_w, + ] = display + + # ------------------------------------------------------------- + # Polygon in progress + # ------------------------------------------------------------- + + if self.polygon_points: + color = class_color_bgr( + self.classes, + self.selected_class, + ) + + pts_win = np.array( + [ + self.source_to_window( + px, + py, + ) + for px, py in self.polygon_points + ], + dtype=np.int32, + ) + + for px, py in pts_win: + cv2.circle( + canvas, + (int(px), int(py)), + 4, + color, + -1, + cv2.LINE_AA, + ) + cv2.circle( + canvas, + (int(px), int(py)), + 6, + (255, 255, 255), + 1, + cv2.LINE_AA, + ) + + if len(pts_win) >= 2: + cv2.polylines( + canvas, + [pts_win], + isClosed=False, + color=color, + thickness=2, + lineType=cv2.LINE_AA, + ) + + # ------------------------------------------------------------- + # Brush cursor + # ------------------------------------------------------------- + + if self.tool == self.TOOL_BRUSH: + mouse_win = self.source_to_window( + self.mouse_src[0], + self.mouse_src[1], + ) + + rr = max( + 2, + int( + round( + self.brush_radius + * self.view.scale + ) + ), + ) + + cv2.circle( + canvas, + mouse_win, + rr, + (255, 255, 255), + 1, + cv2.LINE_AA, + ) + + # ------------------------------------------------------------- + # Header + # ------------------------------------------------------------- + + item = self.samples[self.idx] + complete = self.sample_complete( + item, + ) + + cv2.putText( + canvas, + ( + f"[{self.idx + 1}/{len(self.samples)}] " + f"{item.image_path.name}" + ), + (16, 28), + cv2.FONT_HERSHEY_SIMPLEX, + 0.66, + (245, 245, 245), + 2, + cv2.LINE_AA, + ) + + cv2.putText( + canvas, + ( + f"{'PRONTO' if complete else 'INCOMPLETO'} | " + f"mask={'editada' if self.mask_dirty else 'ok'} | " + f"label={'editado' if self.label_dirty else 'ok'}" + ), + (16, 56), + cv2.FONT_HERSHEY_SIMPLEX, + 0.50, + ( + (70, 220, 80) + if complete + else (70, 180, 255) + ), + 1, + cv2.LINE_AA, + ) + + # ------------------------------------------------------------- + # Right panel + # ------------------------------------------------------------- + + px = MAX_WINDOW_W - PANEL_W + 18 + + cv2.putText( + canvas, + "MASCARA", + (px, 38), + cv2.FONT_HERSHEY_SIMPLEX, + 0.72, + (255, 255, 255), + 2, + cv2.LINE_AA, + ) + + cname = self.classes[ + self.selected_class + ][0] + ccolor = class_color_bgr( + self.classes, + self.selected_class, + ) + + cv2.rectangle( + canvas, + (px, 55), + (px + 28, 83), + ccolor, + -1, + ) + cv2.rectangle( + canvas, + (px, 55), + (px + 28, 83), + (255, 255, 255), + 1, + ) + + cv2.putText( + canvas, + f"Classe: {self.selected_class} {cname}", + (px + 40, 78), + cv2.FONT_HERSHEY_SIMPLEX, + 0.51, + (225, 225, 225), + 1, + cv2.LINE_AA, + ) + + cv2.putText( + canvas, + f"Ferramenta: {self.tool}", + (px, 112), + cv2.FONT_HERSHEY_SIMPLEX, + 0.52, + (220, 220, 220), + 1, + cv2.LINE_AA, + ) + + nav_pct = float( + (self.mask == self.nav_id).mean() + * 100.0 + ) + + cv2.putText( + canvas, + f"Navegavel: {nav_pct:.2f}%", + (px, 142), + cv2.FONT_HERSHEY_SIMPLEX, + 0.52, + (220, 220, 220), + 1, + cv2.LINE_AA, + ) + + try: + group_preview = self.derive_group() + except Exception: + group_preview = "?" + + cv2.putText( + canvas, + f"Grupo: {group_preview}", + (px, 172), + cv2.FONT_HERSHEY_SIMPLEX, + 0.48, + (210, 235, 255), + 1, + cv2.LINE_AA, + ) + + cv2.putText( + canvas, + f"Pincel: {self.brush_radius}px", + (px, 198), + cv2.FONT_HERSHEY_SIMPLEX, + 0.52, + (220, 220, 220), + 1, + cv2.LINE_AA, + ) + + overlay_names = { + self.OVERLAY_NAV_ONLY: "so navegavel", + self.OVERLAY_BOTH: "duas classes", + self.OVERLAY_OFF: "desligado", + } + + cv2.putText( + canvas, + f"Overlay: {overlay_names[self.overlay_mode]}", + (px, 226), + cv2.FONT_HERSHEY_SIMPLEX, + 0.52, + (220, 220, 220), + 1, + cv2.LINE_AA, + ) + + cv2.putText( + canvas, + "P poligono | B pincel | F balde", + (px, 266), + cv2.FONT_HERSHEY_SIMPLEX, + 0.47, + (190, 220, 255), + 1, + cv2.LINE_AA, + ) + + cv2.putText( + canvas, + "C classe | [ ] pincel | O overlay", + (px, 294), + cv2.FONT_HERSHEY_SIMPLEX, + 0.47, + (190, 220, 255), + 1, + cv2.LINE_AA, + ) + + cv2.putText( + canvas, + "Z undo | Y redo | R reset", + (px, 322), + cv2.FONT_HERSHEY_SIMPLEX, + 0.47, + (190, 220, 255), + 1, + cv2.LINE_AA, + ) + + # Status labels. + cv2.putText( + canvas, + "ESTADO DO CORREDOR", + (px, 374), + cv2.FONT_HERSHEY_SIMPLEX, + 0.70, + (255, 255, 255), + 2, + cv2.LINE_AA, + ) + + y = 410 + + for i, state in enumerate( + self.states + ): + key = ( + str(i + 1) + if i < 9 + else "0" + ) + + selected = ( + self.selected_state == i + ) + + color = ( + (80, 255, 120) + if selected + else (205, 205, 205) + ) + + prefix = ">" if selected else " " + + cv2.putText( + canvas, + f"{prefix} [{key}] {state}", + (px, y), + cv2.FONT_HERSHEY_SIMPLEX, + 0.49, + color, + 2 if selected else 1, + cv2.LINE_AA, + ) + + y += 28 + + if self.selected_state is None: + state_txt = "NAO SELECIONADO" + state_color = ( + 80, + 180, + 255, + ) + else: + state_txt = self.states[ + self.selected_state + ] + state_color = ( + 80, + 255, + 120, + ) + + cv2.putText( + canvas, + f"Atual: {state_txt}", + (px, y + 15), + cv2.FONT_HERSHEY_SIMPLEX, + 0.52, + state_color, + 2, + cv2.LINE_AA, + ) + + done = self.completed_count() + + cv2.putText( + canvas, + f"Concluidas: {done}/{len(self.samples)}", + (px, min(MAX_WINDOW_H - 122, y + 58)), + cv2.FONT_HERSHEY_SIMPLEX, + 0.50, + (220, 220, 220), + 1, + cv2.LINE_AA, + ) + + # ------------------------------------------------------------- + # Footer + # ------------------------------------------------------------- + + footer_y = MAX_WINDOW_H - FOOTER_H + 28 + + cv2.putText( + canvas, + ( + "S/SPACE=salvar+proxima | A/D ou setas=navegar | " + "ENTER=fechar poligono | BACKSPACE=cancelar poligono | Q=sair" + ), + (16, footer_y), + cv2.FONT_HERSHEY_SIMPLEX, + 0.49, + (225, 225, 225), + 1, + cv2.LINE_AA, + ) + + cv2.putText( + canvas, + ( + "Pincel: esquerdo=classe selecionada, direito=classe oposta | " + "Mascara bruta salva colorida conforme labelmap no tamanho original" + ), + (16, footer_y + 30), + cv2.FONT_HERSHEY_SIMPLEX, + 0.47, + (165, 205, 255), + 1, + cv2.LINE_AA, + ) + + if ( + self.message + and time.monotonic() + < self.message_until + ): + cv2.putText( + canvas, + self.message[:150], + (16, footer_y + 60), + cv2.FONT_HERSHEY_SIMPLEX, + 0.49, + (80, 255, 160), + 2, + cv2.LINE_AA, + ) + + return canvas + + # ------------------------------------------------------------------------- + # Labels / messages + # ------------------------------------------------------------------------- + + def select_state( + self, + idx: int, + ): + if not ( + 0 <= idx < len(self.states) + ): + return + + self.selected_state = int(idx) + + self.label_dirty = ( + self.loaded_state + != self.selected_state + ) + + self.flash( + f"Estado: {self.states[idx]}" + ) + + def flash( + self, + message: str, + seconds: float = 2.0, + ): + self.message = str(message) + self.message_until = ( + time.monotonic() + + float(seconds) + ) + print(f"[INFO] {message}") + + # ------------------------------------------------------------------------- + # Keys / run + # ------------------------------------------------------------------------- + + def on_key( + self, + k: int, + ) -> bool: + # Quit. + if k in { + 27, + ord("q"), + ord("Q"), + }: + if ( + self.mask_dirty + or self.label_dirty + or self.polygon_points + ): + self.flash( + "Ha alteracoes nao salvas. " + "Pressione Q novamente em ate 2s para sair mesmo assim." + ) + + now = time.monotonic() + + last_q = getattr( + self, + "_last_q", + 0.0, + ) + + if ( + last_q > 0 + and now - last_q <= 2.0 + ): + return False + + self._last_q = now + return True + + return False + + self._last_q = 0.0 + + # Save. + if k in { + ord("s"), + ord("S"), + ord(" "), + }: + self.save_and_next() + return True + + # Tools. + if k in { + ord("p"), + ord("P"), + }: + self.set_tool( + self.TOOL_POLYGON + ) + return True + + if k in { + ord("b"), + ord("B"), + }: + self.set_tool( + self.TOOL_BRUSH + ) + return True + + if k in { + ord("f"), + ord("F"), + }: + self.set_tool( + self.TOOL_FILL + ) + return True + + if k in { + ord("c"), + ord("C"), + }: + self.toggle_class() + return True + + if k == ord("["): + self.brush_radius = max( + 2, + self.brush_radius - 4, + ) + self.flash( + f"Pincel={self.brush_radius}px" + ) + return True + + if k == ord("]"): + self.brush_radius = min( + 300, + self.brush_radius + 4, + ) + self.flash( + f"Pincel={self.brush_radius}px" + ) + return True + + if k in { + ord("o"), + ord("O"), + }: + self.overlay_mode = ( + self.overlay_mode + 1 + ) % 3 + return True + + if k in { + ord("z"), + ord("Z"), + }: + self.undo() + return True + + if k in { + ord("y"), + ord("Y"), + }: + self.redo() + return True + + if k in { + ord("r"), + ord("R"), + }: + self.reset_non_nav() + return True + + # Polygon commit. + if k in { + 10, + 13, + }: + if ( + self.tool + == self.TOOL_POLYGON + ): + self.commit_polygon() + return True + + # Backspace / Delete. + if k in { + 8, + 127, + 3014656, + }: + if self.polygon_points: + self.cancel_polygon() + return True + + # Global labels. + if ( + ord("1") + <= k + <= ord("9") + ): + idx = int(chr(k)) - 1 + self.select_state(idx) + return True + + if k == ord("0"): + self.select_state(9) + return True + + # Navigation ASCII. + if k in { + ord("a"), + ord("A"), + }: + self.navigate(-1) + return True + + if k in { + ord("d"), + ord("D"), + }: + self.navigate(+1) + return True + + # OpenCV Windows arrows via waitKeyEx. + # Common values: + # left=2424832, right=2555904 + if k == 2424832: + self.navigate(-1) + return True + + if k == 2555904: + self.navigate(+1) + return True + + return True + + def run(self): + try: + while not self.finished: + canvas = self.build_canvas() + + cv2.imshow( + WINDOW_NAME, + canvas, + ) + + k = cv2.waitKeyEx(20) + + if k == -1: + continue + + keep = self.on_key( + int(k), + ) + + if not keep: + break + + finally: + cv2.destroyAllWindows() + + +# ============================================================================= +# Main +# ============================================================================= + +def find_archived_bases(original_group_root: Path) -> set[str]: + out: set[str] = set() + if not original_group_root.is_dir(): + return out + + for group_dir in original_group_root.iterdir(): + if not group_dir.is_dir(): + continue + image_dir = group_dir / "images" + if not image_dir.is_dir(): + continue + for path in image_dir.iterdir(): + if path.is_file() and path.suffix.lower() in IMG_EXTS: + out.add(path.stem) + return out + + +def collect_samples( + brutas_dir: Path, + original_group_root: Path, + skip_archived: bool = True, +) -> Tuple[List[SampleItem], int]: + images = list_images(brutas_dir) + archived = find_archived_bases(original_group_root) + + samples: List[SampleItem] = [] + for img in images: + if skip_archived and img.stem in archived: + continue + + # Paths pendentes só existem para manter compatibilidade interna do editor. + # O destino real é resolvido no SAVE depois que o grupo é derivado da máscara. + samples.append( + SampleItem( + image_path=img, + mask_path=Path("__pending__") / "masks" / f"{img.stem}.png", + label_path=Path("__pending__") / "labels" / f"{img.stem}.json", + base=img.stem, + ) + ) + + return samples, len(archived) + + +def main(): + parser = argparse.ArgumentParser( + description=( + "Editor integrado de máscara + estado. Consome dataset/brutas/ " + "e arquiva em dataset/original/group//." + ) + ) + + parser.add_argument("--config", default=str(DEFAULT_CONFIG)) + parser.add_argument("--brutas", default=str(DEFAULT_BRUTAS_DIR)) + parser.add_argument( + "--original-root", + default=str(DEFAULT_ORIGINAL_GROUP_ROOT), + help="Raiz dos grupos finais originais.", + ) + parser.add_argument("--labelmap", default=str(DEFAULT_LABELMAP)) + parser.add_argument( + "--mode", + choices=["move", "copy"], + default="move", + help="move=remove de brutas após commit; copy=preserva a PNG bruta.", + ) + parser.add_argument( + "--include-archived", + action="store_true", + help="Também mostra bases que já existem em original/group (normalmente não use).", + ) + parser.add_argument("--start-index", type=int, default=None) + parser.add_argument("--alpha", type=float, default=DEFAULT_OVERLAY_ALPHA) + parser.add_argument("--brush-radius", type=int, default=DEFAULT_BRUSH_RADIUS) + + args = parser.parse_args() + + config_path = Path(args.config) + if not config_path.is_file(): + raise FileNotFoundError(f"config.json nao encontrado: {config_path}") + + with config_path.open("r", encoding="utf-8") as f: + config = json.load(f) + + states = config.get("label_classes") + if not isinstance(states, list) or not states: + raise RuntimeError("config['label_classes'] obrigatorio.") + if len(states) > 10: + raise RuntimeError("Editor suporta ate 10 estados nas teclas 1..9/0.") + + main_class_name = str(config.get("main_class_name", "navegavel")) + + brutas_dir = Path(args.brutas) + original_group_root = Path(args.original_root) + labelmap_path = Path(args.labelmap) + + classes = load_labelmap_classes(labelmap_path) + nav_id, non_nav_id = resolve_nav_classes(classes, main_class_name) + + samples, archived_before = collect_samples( + brutas_dir, + original_group_root, + skip_archived=not args.include_archived, + ) + + if not samples: + print("[INFO] Nenhuma PNG pendente em brutas. Fila vazia.") + return + + dataset_root = Path("dataset") + + print("=" * 78) + print("Agrobot | Editor de Corredor OAK-D Lite") + print("=" * 78) + print(f"Fila brutas : {brutas_dir.resolve()}") + print(f"Destino original : {original_group_root.resolve()}") + print(f"Modo : {args.mode.upper()}") + print(f"Pendentes : {len(samples)}") + print(f"Já arquivadas : {archived_before}") + print(f"Labelmap : {labelmap_path.resolve()}") + print(f"Navegavel : id={nav_id} name={classes[nav_id][0]}") + print(f"Nao navegavel : id={non_nav_id} name={classes[non_nav_id][0]}") + print("Grupos automáticos: navegavel | naonavegavel | naonavegavel_navegavel") + print("Estados:") + for i, state in enumerate(states): + key = str(i + 1) if i < 9 else "0" + print(f" [{key}] {state}") + print("=" * 78) + + editor = CorridorAnnotationEditor( + dataset_root=dataset_root, + samples=samples, + states=states, + classes=classes, + nav_id=nav_id, + non_nav_id=non_nav_id, + original_group_root=original_group_root, + archive_mode=args.mode, + archived_before=archived_before, + overlay_alpha=args.alpha, + brush_radius=args.brush_radius, + start_index=args.start_index, + ) + + editor.run() + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/oak-d/_2_normalize_corridor.py b/Python/OAK/datasets/oak-d/_2_normalize_corridor.py new file mode 100644 index 000000000..86eb8ebc6 --- /dev/null +++ b/Python/OAK/datasets/oak-d/_2_normalize_corridor.py @@ -0,0 +1,1938 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +_2_normalize_corridor.py +======================== + +Normaliza o dataset ORIGINAL agrupado da OAK-D Lite para a resolução do modelo. + +Executar de dentro da pasta oak-d: + + python _2_normalize_corridor.py + +Entrada: + dataset/original/group/ + navegavel/ + images/.png + masks/.png # RGB colorida conforme labelmap + labels/.json + naonavegavel/ + images/ + masks/ + labels/ + naonavegavel_navegavel/ + images/ + masks/ + labels/ + +Saída, preservando grupos: + dataset/x/group/ + navegavel/{images,masks,labels} + naonavegavel/{images,masks,labels} + naonavegavel_navegavel/{images,masks,labels} + normalize_report.json + +Contrato da saída: + images: PNG RGB/BGR convencional, somente resize INTER_AREA + masks : PNG 1 canal uint8 com IDs de classe, resize INTER_NEAREST + labels: JSON global atualizado e rastreável + +Este script NÃO: + - seleciona amostras; + - faz augmentation offline; + - calcula mean/std; + - cria train/val. + +Pareamento: + - dados novos: image/mask/label usam o mesmo stem; + - legado: sufixos como _Rgb, _mask, _label são reconciliados; + - apenas tripletas completas entram na saída; + - incompletos não abortam por padrão: são registrados no normalize_report; + - poucos pixels RGB fora do labelmap tentam reparo conservador; + - falha de conteúdo vai para dataset/normalize_quarantine e o lote continua; + - por padrão invalida split/norm_stats/report/manifest antigos da mesma resolução. + +mean/std serão calculados DEPOIS do split, somente sobre train. +""" + +from __future__ import annotations + +import argparse +import json +import os +import re +import shutil +from collections import Counter, defaultdict +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from typing import Dict, List, Optional, Sequence, Tuple + +import cv2 +import numpy as np + + +# ============================================================================= +# Defaults +# ============================================================================= + +CONFIG_PATH = Path("config.json") +DATASET_ROOT = Path("dataset") +DEFAULT_ORIGINAL_GROUP_ROOT = DATASET_ROOT / "original" / "group" +DEFAULT_LABELMAP = DATASET_ROOT / "labelmap.txt" +DEFAULT_QUARANTINE_ROOT = DATASET_ROOT / "normalize_quarantine" + +IMG_EXTS = {".png", ".jpg", ".jpeg", ".bmp", ".webp"} +MASK_EXTS = {".png", ".jpg", ".jpeg", ".bmp", ".webp"} +LABEL_EXTS = {".json"} + +VALID_GROUPS = { + "navegavel", + "naonavegavel", + "naonavegavel_navegavel", +} + +REPORT_SCHEMA = "agrobot_corridor_normalize_grouped_v2" +NORMALIZED_LABEL_SCHEMA = "agrobot_corridor_state_normalized_v2" + +FALLBACK_COLORS_RGB = { + 0: (210, 55, 55), + 1: (45, 215, 70), +} + + +# ============================================================================= +# Data +# ============================================================================= + +@dataclass(frozen=True) +class Sample: + group: str + base: str + image: Path + mask: Path + label: Path + + +# ============================================================================= +# Generic helpers +# ============================================================================= + +def now_iso() -> str: + return datetime.now().isoformat(timespec="seconds") + + +def natural_key(text: str): + parts = re.split(r"(\d+)", str(text)) + return [int(p) if p.isdigit() else p.lower() for p in parts] + + +def safe_rel(path: Path, root: Path) -> str: + try: + return str(path.resolve().relative_to(root.resolve())).replace("\\", "/") + except Exception: + return str(path).replace("\\", "/") + + +def atomic_write_json(path: Path, data: dict) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_suffix(path.suffix + ".tmp") + with tmp.open("w", encoding="utf-8") as f: + json.dump(data, f, ensure_ascii=False, indent=2) + os.replace(tmp, path) + + +def atomic_write_png(path: Path, arr: np.ndarray) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_name(path.stem + ".__tmp__.png") + ok = cv2.imwrite(str(tmp), arr, [cv2.IMWRITE_PNG_COMPRESSION, 3]) + if not ok: + try: + tmp.unlink(missing_ok=True) + except Exception: + pass + raise IOError(f"cv2.imwrite retornou False: {tmp}") + os.replace(tmp, path) + + +def map_by_stem(folder: Path, exts: set[str]) -> Dict[str, Path]: + if not folder.is_dir(): + return {} + + result: Dict[str, Path] = {} + for path in sorted(folder.iterdir(), key=lambda p: natural_key(p.name)): + if not path.is_file() or path.suffix.lower() not in exts: + continue + if path.stem in result: + raise RuntimeError( + f"Stem duplicado '{path.stem}' em {folder}: " + f"{result[path.stem].name} e {path.name}" + ) + result[path.stem] = path + return result + + +LEGACY_BASE_SUFFIXES = ( + "_rgb", + "_image", + "_img", + "_frame", + "_mask", + "_masks", + "_seg", + "_segment", + "_segmentacao", + "_label", + "_labels", +) + + +def canonical_base(stem: str) -> str: + """ + Normaliza nomes históricos para casar a mesma amostra. + + Exemplos: + 500_Rgb.png -> 500 + 500_mask.png -> 500 + 500_label.json -> 500 + + O novo pipeline já produz o mesmo stem nas três pastas, então para dados + novos esta função é essencialmente identidade. + """ + out = str(stem) + + changed = True + while changed: + changed = False + low = out.lower() + + for suffix in LEGACY_BASE_SUFFIXES: + if low.endswith(suffix) and len(out) > len(suffix): + out = out[:-len(suffix)] + changed = True + break + + return out + + +def map_by_canonical_base( + folder: Path, + exts: set[str], +) -> Tuple[Dict[str, Path], List[dict], List[dict]]: + """ + Retorna: + canonical -> arquivo + aliases -> casos em que stem original != canonical + collisions -> dois arquivos da MESMA pasta colidiram no mesmo canonical + + Colisão não é resolvida silenciosamente. O canonical fica fora do mapa para + impedir que o normalize escolha uma versão arbitrária. + """ + if not folder.is_dir(): + return {}, [], [] + + buckets: Dict[str, List[Path]] = {} + + for p in sorted(folder.iterdir(), key=lambda x: natural_key(x.name)): + if not p.is_file() or p.suffix.lower() not in exts: + continue + + base = canonical_base(p.stem) + buckets.setdefault(base, []).append(p) + + result: Dict[str, Path] = {} + aliases: List[dict] = [] + collisions: List[dict] = [] + + for base, candidates in sorted( + buckets.items(), + key=lambda kv: natural_key(kv[0]), + ): + if len(candidates) != 1: + collisions.append({ + "base": base, + "files": [p.name for p in candidates], + "folder": str(folder), + }) + continue + + p = candidates[0] + result[base] = p + + if p.stem != base: + aliases.append({ + "base": base, + "original_stem": p.stem, + "file": p.name, + "folder": str(folder), + }) + + return result, aliases, collisions + + +# ============================================================================= +# Labelmap +# ============================================================================= + +def load_labelmap_classes( + path: Path, +) -> Dict[int, Tuple[str, Tuple[int, int, int]]]: + if not path.is_file(): + raise FileNotFoundError(f"labelmap não encontrado: {path}") + + classes: Dict[int, Tuple[str, Tuple[int, int, int]]] = {} + next_id = 0 + + with path.open("r", encoding="utf-8") as f: + for raw in f: + s = raw.strip() + if not s or s.startswith("#"): + continue + + cid: Optional[int] = None + name: Optional[str] = None + color: Optional[Tuple[int, int, int]] = None + left = s.split(":", 1)[0].strip() + + if ":" in s and not left.isdigit(): + name_part, rest = s.split(":", 1) + name = name_part.strip() + color_text = rest.split("::", 1)[0].strip().strip(":") + parts = [p.strip() for p in color_text.split(",") if p.strip()] + if len(parts) >= 3: + try: + color = tuple(int(float(x)) for x in parts[:3]) + except Exception: + color = None + else: + parts = s.replace(",", " ").replace(":", " ").split() + if len(parts) >= 2 and parts[0].isdigit(): + cid = int(parts[0]) + name = parts[1] + if len(parts) >= 5: + try: + color = tuple(int(float(x)) for x in parts[2:5]) + except Exception: + color = None + elif parts: + name = parts[0] + + if not name: + continue + + if name.strip().lower() in {"ignore", "void", "background_ignore"}: + continue + + if cid is None: + cid = next_id + next_id = max(next_id, cid + 1) + + if color is None: + color = FALLBACK_COLORS_RGB.get(cid, (255, 255, 255)) + + classes[int(cid)] = (str(name), tuple(map(int, color))) + + if len(classes) != 2: + raise RuntimeError( + "Normalize especializado em duas classes. " + f"labelmap carregado={classes}" + ) + + return classes + + +def class_color_bgr(classes, cid: int) -> Tuple[int, int, int]: + _name, rgb = classes[int(cid)] + return int(rgb[2]), int(rgb[1]), int(rgb[0]) + + +def resolve_nav_classes(classes, main_class_name: str) -> Tuple[int, int]: + target = str(main_class_name).strip().lower() + + nav = [ + int(cid) + for cid, (name, _rgb) in classes.items() + if str(name).strip().lower() == target + ] + + if len(nav) != 1: + aliases = {"navegavel", "navegável", "nav", "navigable"} + nav = [ + int(cid) + for cid, (name, _rgb) in classes.items() + if str(name).strip().lower() in aliases + ] + + if len(nav) != 1: + raise RuntimeError( + f"Não consegui resolver navegável. main_class_name={main_class_name!r} " + f"classes={classes}" + ) + + nav_id = nav[0] + other = [int(cid) for cid in classes if int(cid) != nav_id] + if len(other) != 1: + raise RuntimeError(f"Classe não navegável ambígua: {classes}") + return nav_id, other[0] + + +# ============================================================================= +# Mask contract +# ============================================================================= + +def _exact_color_decode( + mask_bgr: np.ndarray, + classes, +) -> Tuple[np.ndarray, np.ndarray]: + """Converte somente cores EXATAS do labelmap. Retorna ids + unknown mask.""" + h, w = mask_bgr.shape[:2] + ids = np.full((h, w), 255, dtype=np.uint8) + + for cid in sorted(classes): + bgr = np.array( + class_color_bgr(classes, cid), + dtype=np.uint8, + ) + match = np.all( + mask_bgr[:, :, :3] == bgr[None, None, :], + axis=2, + ) + ids[match] = int(cid) + + return ids, ids == 255 + + +def _repair_unknown_spatial( + ids: np.ndarray, + classes, + max_iters: int = 6, +) -> Tuple[np.ndarray, int]: + """ + Recupera pixels desconhecidos quando a vizinhança é fortemente consistente. + + Não usa interpolação na máscara. O pixel só herda uma classe se houver + maioria local clara entre os 8 vizinhos já conhecidos. + """ + out = ids.copy() + repaired = 0 + class_ids = [int(cid) for cid in sorted(classes)] + + for _ in range(max(1, int(max_iters))): + unknown = out == 255 + if not np.any(unknown): + break + + padded = np.pad( + out, + 1, + mode="constant", + constant_values=255, + ) + + neighbors = [ + padded[0:-2, 0:-2], + padded[0:-2, 1:-1], + padded[0:-2, 2:], + padded[1:-1, 0:-2], + padded[1:-1, 2:], + padded[2:, 0:-2], + padded[2:, 1:-1], + padded[2:, 2:], + ] + + votes = np.stack( + [ + sum((n == cid).astype(np.uint8) for n in neighbors) + for cid in class_ids + ], + axis=0, + ) + + order = np.argsort(votes, axis=0) + top_idx = order[-1] + top_votes = np.take_along_axis( + votes, + top_idx[None, ...], + axis=0, + )[0] + + if len(class_ids) > 1: + second_idx = order[-2] + second_votes = np.take_along_axis( + votes, + second_idx[None, ...], + axis=0, + )[0] + else: + second_votes = np.zeros_like(top_votes) + + margin = top_votes.astype(np.int16) - second_votes.astype(np.int16) + + # Conservador: + # pelo menos 4 vizinhos concordam e há margem >= 3 sobre a outra classe. + accept = ( + unknown + & (top_votes >= 4) + & (margin >= 3) + ) + + if not np.any(accept): + break + + top_class = np.take( + np.asarray(class_ids, dtype=np.uint8), + top_idx, + ) + + out[accept] = top_class[accept] + repaired += int(accept.sum()) + + return out, repaired + + +def _repair_unknown_color( + ids: np.ndarray, + mask_bgr: np.ndarray, + classes, + max_distance: float = 55.0, + min_margin: float = 20.0, +) -> Tuple[np.ndarray, int]: + """ + Segunda chance por distância de cor, mas somente quando a melhor classe + está claramente mais próxima que a segunda. + + Isso NÃO é nearest-color irrestrito. + """ + out = ids.copy() + unknown = out == 255 + + if not np.any(unknown): + return out, 0 + + class_ids = [int(cid) for cid in sorted(classes)] + colors = np.asarray( + [ + class_color_bgr(classes, cid) + for cid in class_ids + ], + dtype=np.float32, + ) + + pixels = mask_bgr[:, :, :3][unknown].astype(np.float32) + + # [N, K] + dist = np.linalg.norm( + pixels[:, None, :] - colors[None, :, :], + axis=2, + ) + + nearest_idx = np.argmin(dist, axis=1) + nearest_dist = dist[ + np.arange(dist.shape[0]), + nearest_idx, + ] + + if dist.shape[1] >= 2: + sorted_dist = np.sort(dist, axis=1) + margin = sorted_dist[:, 1] - sorted_dist[:, 0] + else: + margin = np.full_like( + nearest_dist, + np.inf, + ) + + accept = ( + (nearest_dist <= float(max_distance)) + & (margin >= float(min_margin)) + ) + + if not np.any(accept): + return out, 0 + + unknown_yx = np.argwhere(unknown) + accepted_yx = unknown_yx[accept] + accepted_ids = np.asarray( + class_ids, + dtype=np.uint8, + )[nearest_idx[accept]] + + out[ + accepted_yx[:, 0], + accepted_yx[:, 1], + ] = accepted_ids + + return out, int(accept.sum()) + + +def _repair_unknown_nearest_class( + ids: np.ndarray, + classes, + non_nav_id: Optional[int] = None, +) -> Tuple[np.ndarray, int, int]: + """ + Fallback geométrico para máscaras categóricas legadas. + + Depois das tentativas conservadoras, pixels ainda desconhecidos são + atribuídos à classe VÁLIDA espacialmente mais próxima. + + Caso duas classes estejam à mesma distância, o empate vai para + `non_nav_id`, quando disponível. Isso é deliberadamente conservador: + uma fronteira antiga/anti-aliased de 1 px não deve virar falso navegável. + + Retorna: + ids_reparados + quantidade_reparada + quantidade_de_empates_resolvidos_para_non_nav + """ + out = ids.copy() + unknown = out == 255 + + n_unknown = int(unknown.sum()) + + if n_unknown == 0: + return out, 0, 0 + + present_class_ids = [ + int(cid) + for cid in sorted(classes) + if np.any(out == int(cid)) + ] + + if not present_class_ids: + return out, 0, 0 + + # Máscara pura com poucos artefatos: só existe uma classe válida. + if len(present_class_ids) == 1: + out[unknown] = int(present_class_ids[0]) + return out, n_unknown, 0 + + distance_maps = [] + + for cid in present_class_ids: + # distanceTransform mede distância ao zero mais próximo. + # Portanto pixels da classe atual viram zero. + src = np.ones( + out.shape, + dtype=np.uint8, + ) + src[out == int(cid)] = 0 + + dist = cv2.distanceTransform( + src, + cv2.DIST_L2, + 3, + ) + + distance_maps.append( + dist + ) + + dists = np.stack( + distance_maps, + axis=0, + ) + + ys, xs = np.where( + unknown + ) + + d_unknown = dists[ + :, + ys, + xs, + ] + + min_dist = np.min( + d_unknown, + axis=0, + ) + + # Tolerância pequena porque distanceTransform usa aproximação. + tie_eps = 1e-3 + + assignments = np.empty( + n_unknown, + dtype=np.uint8, + ) + + tie_to_non_nav = 0 + + non_nav_available = ( + non_nav_id is not None + and int(non_nav_id) + in present_class_ids + ) + + class_ids_arr = np.asarray( + present_class_ids, + dtype=np.uint8, + ) + + for i in range(n_unknown): + tied = np.where( + np.abs( + d_unknown[:, i] + - min_dist[i] + ) + <= tie_eps + )[0] + + if len(tied) == 1: + assignments[i] = class_ids_arr[ + tied[0] + ] + continue + + if non_nav_available: + non_nav_pos = present_class_ids.index( + int(non_nav_id) + ) + + if non_nav_pos in tied: + assignments[i] = int( + non_nav_id + ) + tie_to_non_nav += 1 + continue + + # Fallback determinístico se non_nav não estiver entre os empatados. + tied_ids = class_ids_arr[ + tied + ] + assignments[i] = int( + np.min( + tied_ids + ) + ) + + out[ + ys, + xs, + ] = assignments + + return ( + out, + n_unknown, + tie_to_non_nav, + ) + + +def colored_mask_to_ids( + mask_bgr: np.ndarray, + classes, + *, + non_nav_id: Optional[int] = None, + repair_enabled: bool = True, + max_unknown_pixels: int = 4096, + max_unknown_fraction: float = 0.002, + color_max_distance: float = 55.0, + color_min_margin: float = 20.0, +) -> Tuple[np.ndarray, dict]: + """ + RGB/BGR -> IDs com recuperação CONSERVADORA. + + Ordem: + 1) cores exatas do labelmap; + 2) consenso espacial; + 3) distância de cor com margem forte; + 4) fallback geométrico para a classe válida espacialmente mais próxima; + empate -> não navegável; + 5) quarentena somente se o volume de pixels inválidos já exceder + os limites de reparo ou se algo realmente inconsistente permanecer. + """ + ids, unknown = _exact_color_decode( + mask_bgr, + classes, + ) + + unknown_initial = int(unknown.sum()) + total_pixels = int(ids.size) + + info = { + "enabled": bool(repair_enabled), + "unknown_initial": unknown_initial, + "unknown_fraction_initial": ( + float(unknown_initial / max(1, total_pixels)) + ), + "repaired_spatial": 0, + "repaired_color": 0, + "repaired_nearest": 0, + "nearest_ties_to_non_nav": 0, + "unknown_final": unknown_initial, + "repaired": False, + } + + if unknown_initial == 0: + return ids, info + + pixels = mask_bgr[:, :, :3][unknown] + uniq = np.unique( + pixels.reshape(-1, 3), + axis=0, + ) + + if not repair_enabled: + raise ValueError( + f"Máscara possui {unknown_initial} pixels com cores fora do labelmap. " + f"Primeiras cores BGR={uniq[:10].tolist()}" + ) + + fraction = unknown_initial / max(1, total_pixels) + + if ( + unknown_initial > int(max_unknown_pixels) + or fraction > float(max_unknown_fraction) + ): + raise ValueError( + f"Máscara possui {unknown_initial} pixels desconhecidos " + f"({fraction * 100.0:.5f}%), acima do limite de reparo " + f"(pixels<={int(max_unknown_pixels)}, " + f"fração<={float(max_unknown_fraction) * 100.0:.5f}%). " + f"Primeiras cores BGR={uniq[:10].tolist()}" + ) + + ids, n_spatial = _repair_unknown_spatial( + ids, + classes, + ) + info["repaired_spatial"] = int(n_spatial) + + ids, n_color = _repair_unknown_color( + ids, + mask_bgr, + classes, + max_distance=float(color_max_distance), + min_margin=float(color_min_margin), + ) + info["repaired_color"] = int(n_color) + + # Legado: anti-alias/ruído de 1 px frequentemente cai EXATAMENTE + # na fronteira das duas classes. Nesse caso maioria local e distância + # de cor podem ficar empatadas. Resolve pela geometria da máscara. + ids, n_nearest, n_ties_non_nav = _repair_unknown_nearest_class( + ids, + classes, + non_nav_id=non_nav_id, + ) + info["repaired_nearest"] = int(n_nearest) + info["nearest_ties_to_non_nav"] = int(n_ties_non_nav) + + remaining = ids == 255 + unknown_final = int(remaining.sum()) + + info["unknown_final"] = unknown_final + info["repaired"] = ( + unknown_initial > 0 + and unknown_final == 0 + ) + + if unknown_final: + pixels_final = mask_bgr[:, :, :3][remaining] + uniq_final = np.unique( + pixels_final.reshape(-1, 3), + axis=0, + ) + + raise ValueError( + f"Máscara tinha {unknown_initial} pixels fora do labelmap; " + f"reparo resolveu " + f"{int(n_spatial) + int(n_color) + int(n_nearest)}, " + f"mas restaram {unknown_final} " + f"pixels ambíguos. Primeiras cores BGR restantes=" + f"{uniq_final[:10].tolist()}" + ) + + return ids, info + + +def load_mask_ids( + path: Path, + classes, + expected_hw: Tuple[int, int], + *, + non_nav_id: Optional[int] = None, + repair_enabled: bool = True, + max_unknown_pixels: int = 4096, + max_unknown_fraction: float = 0.002, + color_max_distance: float = 55.0, + color_min_margin: float = 20.0, +): + raw = cv2.imread( + str(path), + cv2.IMREAD_UNCHANGED, + ) + + if raw is None: + raise FileNotFoundError( + f"Não consegui abrir mask: {path}" + ) + + if raw.shape[:2] != expected_hw: + raise ValueError( + f"Mask {path.name} shape={raw.shape[:2]} != imagem {expected_hw}" + ) + + repair_info = { + "enabled": bool(repair_enabled), + "unknown_initial": 0, + "unknown_fraction_initial": 0.0, + "repaired_spatial": 0, + "repaired_color": 0, + "repaired_nearest": 0, + "nearest_ties_to_non_nav": 0, + "unknown_final": 0, + "repaired": False, + } + + if raw.ndim == 2: + ids = raw.astype( + np.uint8, + copy=False, + ) + contract = "ids" + + elif raw.ndim == 3 and raw.shape[2] >= 3: + ids, repair_info = colored_mask_to_ids( + raw[:, :, :3], + classes, + non_nav_id=non_nav_id, + repair_enabled=repair_enabled, + max_unknown_pixels=max_unknown_pixels, + max_unknown_fraction=max_unknown_fraction, + color_max_distance=color_max_distance, + color_min_margin=color_min_margin, + ) + + contract = ( + "rgb_color_repaired" + if repair_info["repaired"] + else "rgb_color" + ) + + else: + raise ValueError( + f"Mask inválida {path.name}: shape={raw.shape}" + ) + + present = set( + int(x) + for x in np.unique(ids) + ) + allowed = set( + int(cid) + for cid in classes + ) + + if not present.issubset(allowed): + raise ValueError( + f"Mask {path.name} IDs inválidos=" + f"{sorted(present - allowed)}" + ) + + return ids.copy(), contract, repair_info + + +def build_unknown_mask_diagnostic( + mask_path: Path, + classes, +) -> Optional[np.ndarray]: + """ + Gera uma cópia visual da máscara com pixels de cor desconhecida em magenta. + Útil para revisão manual da quarentena. + """ + raw = cv2.imread( + str(mask_path), + cv2.IMREAD_UNCHANGED, + ) + + if ( + raw is None + or raw.ndim != 3 + or raw.shape[2] < 3 + ): + return None + + vis = raw[:, :, :3].copy() + _ids, unknown = _exact_color_decode( + vis, + classes, + ) + + if not np.any(unknown): + return None + + # BGR magenta. + vis[unknown] = np.array( + [255, 0, 255], + dtype=np.uint8, + ) + + return vis + + +def derive_group(mask_ids: np.ndarray, nav_id: int, non_nav_id: int) -> str: + present = set(int(x) for x in np.unique(mask_ids)) + + if present == {non_nav_id}: + return "naonavegavel" + if present == {nav_id}: + return "navegavel" + if present == {nav_id, non_nav_id}: + return "naonavegavel_navegavel" + + raise ValueError(f"Composição inesperada da mask: IDs={sorted(present)}") + + +# ============================================================================= +# Dataset grouped +# ============================================================================= + +def collect_group_samples( + original_group_root: Path, + strict: bool = False, +) -> Tuple[List[Sample], dict]: + """ + Coleta SOMENTE tripletas completas. + + Importante: + - nomes legados são reconciliados por canonical_base(); + - item incompleto é ignorado por padrão e registrado no relatório; + - colisão ambígua de nomes também é ignorada/registrada; + - --strict-input transforma esses casos em erro. + + Portanto: + 500_Rgb.png + 500.png + 500.json + é uma tripleta válida com base canônica "500". + """ + if not original_group_root.is_dir(): + raise FileNotFoundError( + f"Raiz original agrupada não encontrada: {original_group_root}" + ) + + samples: List[Sample] = [] + group_report = {} + + all_incomplete: List[dict] = [] + all_collisions: List[dict] = [] + all_aliases: List[dict] = [] + + group_dirs = sorted( + [p for p in original_group_root.iterdir() if p.is_dir()], + key=lambda p: natural_key(p.name), + ) + + for group_dir in group_dirs: + group = group_dir.name + + images, image_aliases, image_collisions = map_by_canonical_base( + group_dir / "images", + IMG_EXTS, + ) + masks, mask_aliases, mask_collisions = map_by_canonical_base( + group_dir / "masks", + MASK_EXTS, + ) + labels, label_aliases, label_collisions = map_by_canonical_base( + group_dir / "labels", + LABEL_EXTS, + ) + + aliases = image_aliases + mask_aliases + label_aliases + + collisions = [] + for modality, rows in ( + ("image", image_collisions), + ("mask", mask_collisions), + ("label", label_collisions), + ): + for row in rows: + row = dict(row) + row["group"] = group + row["modality"] = modality + collisions.append(row) + + for row in aliases: + row["group"] = group + + all_aliases.extend(aliases) + all_collisions.extend(collisions) + + # Canonicals com colisão são conhecidos, mas foram removidos do map. + collision_bases_by_modality = { + "image": {r["base"] for r in image_collisions}, + "mask": {r["base"] for r in mask_collisions}, + "label": {r["base"] for r in label_collisions}, + } + + all_bases = sorted( + set(images) + | set(masks) + | set(labels) + | collision_bases_by_modality["image"] + | collision_bases_by_modality["mask"] + | collision_bases_by_modality["label"], + key=natural_key, + ) + + complete = 0 + incomplete = [] + + for base in all_bases: + missing = [] + ambiguous = [] + + if base in collision_bases_by_modality["image"]: + ambiguous.append("image") + elif base not in images: + missing.append("image") + + if base in collision_bases_by_modality["mask"]: + ambiguous.append("mask") + elif base not in masks: + missing.append("mask") + + if base in collision_bases_by_modality["label"]: + ambiguous.append("label") + elif base not in labels: + missing.append("label") + + if missing or ambiguous: + row = { + "group": group, + "base": base, + "missing": missing, + "ambiguous": ambiguous, + } + incomplete.append(row) + all_incomplete.append(row) + continue + + samples.append( + Sample( + group=group, + base=base, + image=images[base], + mask=masks[base], + label=labels[base], + ) + ) + complete += 1 + + group_report[group] = { + "images_canonical": len(images), + "masks_canonical": len(masks), + "labels_canonical": len(labels), + "complete": complete, + "incomplete_count": len(incomplete), + "incomplete": incomplete, + "legacy_alias_count": len(aliases), + "legacy_aliases": aliases, + "collision_count": len(collisions), + "collisions": collisions, + } + + if strict and (all_incomplete or all_collisions): + lines = [] + + for x in all_incomplete[:20]: + details = [] + if x.get("missing"): + details.append( + "faltando " + ",".join(x["missing"]) + ) + if x.get("ambiguous"): + details.append( + "ambíguo " + ",".join(x["ambiguous"]) + ) + + lines.append( + f"{x['group']}/{x['base']}: " + "; ".join(details) + ) + + raise RuntimeError( + "Dataset original possui itens não utilizáveis no pareamento " + f"(incompletos={len(all_incomplete)}, colisões={len(all_collisions)}).\n" + + "\n".join(lines) + ) + + return samples, { + "groups": group_report, + "complete_count": len(samples), + "incomplete_count": len(all_incomplete), + "collision_count": len(all_collisions), + "legacy_alias_count": len(all_aliases), + "incomplete": all_incomplete, + "collisions": all_collisions, + "legacy_aliases": all_aliases, + } + + +def validate_label(label_path: Path, states: Sequence[str], expected_group: str) -> dict: + try: + with label_path.open("r", encoding="utf-8") as f: + data = json.load(f) + except Exception as exc: + raise ValueError(f"Label JSON inválido {label_path}: {exc}") from exc + + if not isinstance(data, dict): + raise ValueError(f"Label não é objeto JSON: {label_path}") + + if "label_id" not in data: + raise ValueError(f"Label sem label_id: {label_path}") + + lid = int(data["label_id"]) + if not 0 <= lid < len(states): + raise ValueError( + f"label_id={lid} fora de [0,{len(states)-1}] em {label_path}" + ) + + expected_state = str(states[lid]) + got_state = data.get("estado_corredor") + if got_state is not None and str(got_state) != expected_state: + raise ValueError( + f"Label inconsistente {label_path}: id={lid} => {expected_state!r}, " + f"estado_corredor={got_state!r}" + ) + + got_group = data.get("group") + if got_group is not None and str(got_group) != expected_group: + raise ValueError( + f"Label {label_path.name} diz group={got_group!r}, " + f"mas está em {expected_group!r}" + ) + + return data + + +# ============================================================================= +# Downstream invalidation +# ============================================================================= + +def clear_downstream_artifacts( + resolution_root: Path, +) -> List[str]: + """ + Reexecutar NORMALIZE invalida tudo que depende dele. + + Remove: + dataset/x/split/ + dataset/x/norm_stats.json + dataset/x/split_report.json + dataset/x/split_manifest.csv + + A fonte normalizada group/ será limpa separadamente pelo próprio normalize. + """ + removed: List[str] = [] + + split_root = resolution_root / "split" + if split_root.exists(): + shutil.rmtree(split_root) + removed.append(str(split_root)) + + for name in ( + "norm_stats.json", + "split_report.json", + "split_manifest.csv", + ): + p = resolution_root / name + if p.exists(): + p.unlink() + removed.append(str(p)) + + return removed + + +# ============================================================================= +# Quarantine +# ============================================================================= + +def prepare_quarantine_root( + root: Path, + clean: bool, +) -> None: + if clean and root.exists(): + shutil.rmtree(root) + + root.mkdir( + parents=True, + exist_ok=True, + ) + + +def quarantine_sample( + sample: Sample, + *, + quarantine_root: Path, + classes, + error: Exception, +) -> dict: + """ + Copia a tripleta ORIGINAL para revisão manual. + + Estrutura: + normalize_quarantine// + images/ + masks/ + labels/ + diagnostics/ + errors/ + """ + group_root = ( + quarantine_root + / sample.group + ) + + dirs = { + "images": group_root / "images", + "masks": group_root / "masks", + "labels": group_root / "labels", + "diagnostics": group_root / "diagnostics", + "errors": group_root / "errors", + } + + for folder in dirs.values(): + folder.mkdir( + parents=True, + exist_ok=True, + ) + + out_image = ( + dirs["images"] + / f"{sample.base}{sample.image.suffix.lower()}" + ) + out_mask = ( + dirs["masks"] + / f"{sample.base}{sample.mask.suffix.lower()}" + ) + out_label = ( + dirs["labels"] + / f"{sample.base}{sample.label.suffix.lower()}" + ) + + shutil.copy2( + sample.image, + out_image, + ) + shutil.copy2( + sample.mask, + out_mask, + ) + shutil.copy2( + sample.label, + out_label, + ) + + diagnostic_path = None + + try: + diagnostic = build_unknown_mask_diagnostic( + sample.mask, + classes, + ) + + if diagnostic is not None: + diagnostic_path = ( + dirs["diagnostics"] + / f"{sample.base}_unknown.png" + ) + atomic_write_png( + diagnostic_path, + diagnostic, + ) + except Exception: + diagnostic_path = None + + error_path = ( + dirs["errors"] + / f"{sample.base}.json" + ) + + payload = { + "schema": "agrobot_normalize_quarantine_v1", + "created_at": now_iso(), + "group": sample.group, + "base": sample.base, + "error_type": type(error).__name__, + "error": str(error), + "source_image": str(sample.image), + "source_mask": str(sample.mask), + "source_label": str(sample.label), + "quarantine_image": str(out_image), + "quarantine_mask": str(out_mask), + "quarantine_label": str(out_label), + "diagnostic_unknown_pixels": ( + str(diagnostic_path) + if diagnostic_path is not None + else None + ), + } + + atomic_write_json( + error_path, + payload, + ) + + return payload + + +# ============================================================================= +# Processing +# ============================================================================= + +def check_aspect_ratio( + src_wh: Tuple[int, int], + dst_wh: Tuple[int, int], + tolerance: float = 1e-3, +): + sw, sh = src_wh + dw, dh = dst_wh + src_ar = sw / float(sh) + dst_ar = dw / float(dh) + rel = abs(src_ar - dst_ar) / max(dst_ar, 1e-9) + if rel > tolerance: + raise ValueError( + f"Aspect ratio incompatível: origem={sw}x{sh} ({src_ar:.6f}) vs " + f"destino={dw}x{dh} ({dst_ar:.6f}). " + "Não será feito crop/letterbox silencioso." + ) + + +def prepare_output_root(output_group_root: Path, clean: bool): + if clean and output_group_root.exists(): + shutil.rmtree(output_group_root) + output_group_root.mkdir(parents=True, exist_ok=True) + + +def process_sample( + sample: Sample, + *, + target_wh: Tuple[int, int], + classes, + states: Sequence[str], + nav_id: int, + non_nav_id: int, + dataset_root: Path, + output_root: Path, + repair_enabled: bool, + repair_max_unknown_pixels: int, + repair_max_unknown_fraction: float, + repair_color_max_distance: float, + repair_color_min_margin: float, +) -> dict: + target_w, target_h = target_wh + + img = cv2.imread(str(sample.image), cv2.IMREAD_COLOR) + if img is None: + raise FileNotFoundError(f"Não consegui abrir imagem: {sample.image}") + + src_h, src_w = img.shape[:2] + check_aspect_ratio((src_w, src_h), target_wh) + + mask_ids, source_mask_contract, mask_repair = load_mask_ids( + sample.mask, + classes, + expected_hw=(src_h, src_w), + non_nav_id=non_nav_id, + repair_enabled=repair_enabled, + max_unknown_pixels=repair_max_unknown_pixels, + max_unknown_fraction=repair_max_unknown_fraction, + color_max_distance=repair_color_max_distance, + color_min_margin=repair_color_min_margin, + ) + + derived_group = derive_group(mask_ids, nav_id, non_nav_id) + if sample.group != derived_group: + raise ValueError( + f"Grupo divergente em {sample.base}: pasta={sample.group!r}, " + f"mask implica={derived_group!r}" + ) + + label = validate_label(sample.label, states, expected_group=sample.group) + label_id = int(label["label_id"]) + estado = str(states[label_id]) + + group_root = output_root / "group" / sample.group + out_images = group_root / "images" + out_masks = group_root / "masks" + out_labels = group_root / "labels" + + out_img = out_images / f"{sample.base}.png" + out_mask = out_masks / f"{sample.base}.png" + out_label = out_labels / f"{sample.base}.json" + + img_resized = cv2.resize( + img, + (target_w, target_h), + interpolation=cv2.INTER_AREA, + ) + atomic_write_png(out_img, img_resized) + + mask_resized = cv2.resize( + mask_ids, + (target_w, target_h), + interpolation=cv2.INTER_NEAREST, + ).astype(np.uint8, copy=False) + + resized_group = derive_group(mask_resized, nav_id, non_nav_id) + if resized_group != sample.group: + # Em teoria nearest pode eliminar uma classe presente em pouquíssimos pixels. + # Isso é informação importante, não deve passar silenciosamente. + raise ValueError( + f"Após resize, composição mudou de grupo: {sample.group} -> {resized_group}. " + "Revise esta máscara; a classe minoritária pode ser pequena demais." + ) + + atomic_write_png(out_mask, mask_resized) + + nav_pct = float((mask_resized == nav_id).mean() * 100.0) + + normalized = dict(label) + normalized.update({ + "schema": NORMALIZED_LABEL_SCHEMA, + "group": sample.group, + "estado_corredor": estado, + "label_id": label_id, + "states": list(states), + "base": sample.base, + "normalized": True, + "normalized_at": now_iso(), + "normalized_resolution_wh": [target_w, target_h], + "image": safe_rel(out_img, output_root), + "mask": safe_rel(out_mask, output_root), + "label": safe_rel(out_label, output_root), + "source_image": safe_rel(sample.image, dataset_root), + "source_mask": safe_rel(sample.mask, dataset_root), + "source_label": safe_rel(sample.label, dataset_root), + "source_mask_contract": source_mask_contract, + "mask_repair": mask_repair, + "normalized_mask_contract": "uint8_class_ids", + "mask_classes": { + str(cid): { + "name": classes[cid][0], + "rgb": list(classes[cid][1]), + } + for cid in sorted(classes) + }, + "nav_class_id": int(nav_id), + "non_nav_class_id": int(non_nav_id), + "navegavel_pct_normalized": nav_pct, + }) + + atomic_write_json(out_label, normalized) + + return { + "group": sample.group, + "base": sample.base, + "source_wh": [src_w, src_h], + "target_wh": [target_w, target_h], + "source_mask_contract": source_mask_contract, + "mask_repair": mask_repair, + "label_id": label_id, + "estado_corredor": estado, + "navegavel_pct": nav_pct, + } + + +# ============================================================================= +# Main +# ============================================================================= + +def main(): + ap = argparse.ArgumentParser( + description=( + "Normaliza dataset/original/group da OAK-D Lite preservando grupos." + ) + ) + ap.add_argument("--config", default=str(CONFIG_PATH)) + ap.add_argument("--original-root", default=str(DEFAULT_ORIGINAL_GROUP_ROOT)) + ap.add_argument("--labelmap", default=str(DEFAULT_LABELMAP)) + ap.add_argument( + "--strict-input", + action="store_true", + help=( + "Abortar se houver item incompleto ou colisão de nomes. " + "Default: processa apenas tripletas completas e registra o restante." + ), + ) + ap.add_argument( + "--allow-incomplete", + action="store_true", + help=( + "Compatibilidade com versões anteriores. Hoje já é o comportamento " + "default; mantido para não quebrar comandos antigos." + ), + ) + ap.add_argument( + "--continue-on-error", + action="store_true", + help=( + "Compatibilidade: hoje continuar após erro de conteúdo já é o default." + ), + ) + ap.add_argument( + "--fail-fast-content", + action="store_true", + help=( + "Abortar no primeiro erro de conteúdo de uma tripleta completa. " + "Default: envia a tripleta para quarentena e continua." + ), + ) + ap.add_argument( + "--quarantine-root", + default=str(DEFAULT_QUARANTINE_ROOT), + help=( + "Pasta para tripletas que falharem durante normalização." + ), + ) + ap.add_argument( + "--no-mask-repair", + action="store_true", + help=( + "Desativa recuperação conservadora de poucos pixels RGB fora do labelmap." + ), + ) + ap.add_argument( + "--repair-max-unknown-pixels", + type=int, + default=4096, + help="Máximo absoluto de pixels desconhecidos elegíveis a reparo.", + ) + ap.add_argument( + "--repair-max-unknown-fraction", + type=float, + default=0.002, + help=( + "Máxima fração de pixels desconhecidos elegível a reparo. " + "0.002 = 0.2%% da máscara." + ), + ) + ap.add_argument( + "--repair-color-max-distance", + type=float, + default=55.0, + help="Distância BGR máxima para segunda chance por cor.", + ) + ap.add_argument( + "--repair-color-min-margin", + type=float, + default=20.0, + help="Margem mínima entre melhor e segunda melhor cor para aceitar reparo.", + ) + ap.add_argument( + "--no-clean", + action="store_true", + help=( + "Não limpa a saída normalizada anterior e não invalida artefatos " + "downstream. Use somente quando souber exatamente o que está fazendo." + ), + ) + args = ap.parse_args() + + config_path = Path(args.config) + if not config_path.is_file(): + raise FileNotFoundError(f"config não encontrado: {config_path}") + + with config_path.open("r", encoding="utf-8") as f: + config = json.load(f) + + resolution = config.get("resolucao") + if not isinstance(resolution, list) or len(resolution) != 2: + raise RuntimeError("config['resolucao'] deve ser [W,H].") + + target_w, target_h = int(resolution[0]), int(resolution[1]) + if target_w <= 0 or target_h <= 0: + raise RuntimeError(f"Resolução inválida: {resolution}") + + states = config.get("label_classes") + if not isinstance(states, list) or not states: + raise RuntimeError("config['label_classes'] obrigatório.") + + classes = load_labelmap_classes(Path(args.labelmap)) + nav_id, non_nav_id = resolve_nav_classes( + classes, + str(config.get("main_class_name", "navegavel")), + ) + + original_group_root = Path(args.original_root) + + strict_input = bool(args.strict_input) and not bool(args.allow_incomplete) + + # Pipeline legado grande: por padrão erro de conteúdo vai para quarentena + # e o normalize continua. --fail-fast-content restaura o modo rígido. + fail_fast_processing = bool(args.fail_fast_content) + + repair_enabled = not bool(args.no_mask_repair) + + if args.repair_max_unknown_pixels < 0: + raise ValueError("--repair-max-unknown-pixels deve ser >= 0") + + if not (0.0 <= args.repair_max_unknown_fraction <= 1.0): + raise ValueError("--repair-max-unknown-fraction deve ficar entre 0 e 1") + + quarantine_root = Path(args.quarantine_root) + + samples, collect_report = collect_group_samples( + original_group_root, + strict=strict_input, + ) + + if not samples: + raise RuntimeError(f"Nenhuma amostra completa em {original_group_root}") + + output_root = DATASET_ROOT / f"{target_w}x{target_h}" + output_group_root = output_root / "group" + + removed_downstream = [] + + if not args.no_clean: + removed_downstream = clear_downstream_artifacts( + output_root + ) + + prepare_output_root( + output_group_root, + clean=not args.no_clean, + ) + prepare_quarantine_root( + quarantine_root, + clean=not args.no_clean, + ) + + print("=" * 82) + print("OAK-D Lite | Normalize Corridor Dataset por grupos") + print("=" * 82) + print(f"Entrada : {original_group_root.resolve()}") + print(f"Saída : {output_root.resolve()}") + print(f"Resolução : {target_w}x{target_h}") + print(f"Tripletas válidas : {len(samples)}") + print(f"Incompletas skip : {collect_report['incomplete_count']}") + print(f"Aliases legados : {collect_report['legacy_alias_count']}") + print(f"Colisões skip : {collect_report['collision_count']}") + print(f"Navegável : id={nav_id} {classes[nav_id][0]}") + print(f"Não navegável : id={non_nav_id} {classes[non_nav_id][0]}") + print("Augmentation : NÃO, fica no trainer") + print("Mean/std : NÃO, será calculado depois do split apenas no TRAIN") + print( + f"Mask repair : {'ON' if repair_enabled else 'OFF'} " + f"(max_px={args.repair_max_unknown_pixels}, " + f"max_frac={args.repair_max_unknown_fraction:.6f})" + ) + print(f"Quarentena : {quarantine_root.resolve()}") + print( + f"Erro conteúdo : " + f"{'ABORTA' if fail_fast_processing else 'QUARENTENA + CONTINUA'}" + ) + print( + f"Clear downstream : " + f"{'OFF (--no-clean)' if args.no_clean else 'ON'}" + ) + + if removed_downstream: + print("[CLEAN] Artefatos downstream antigos removidos:") + for removed_path in removed_downstream: + print(f" - {removed_path}") + print("=" * 82) + + if collect_report["incomplete_count"]: + print("[PAIRING] Itens incompletos ignorados (primeiros 12):") + for row in collect_report["incomplete"][:12]: + details = [] + if row.get("missing"): + details.append("faltando=" + ",".join(row["missing"])) + if row.get("ambiguous"): + details.append("ambíguo=" + ",".join(row["ambiguous"])) + print( + f" - {row['group']}/{row['base']} | " + + " | ".join(details) + ) + + if collect_report["collision_count"]: + print("[PAIRING] Colisões de nome ignoradas (primeiras 8):") + for row in collect_report["collisions"][:8]: + print( + f" - {row['group']}/{row['base']} " + f"[{row['modality']}] -> {row['files']}" + ) + + if collect_report["legacy_alias_count"]: + print( + f"[PAIRING] {collect_report['legacy_alias_count']} nome(s) legado(s) " + "foram reconciliados por base canônica." + ) + + print() + + processed = [] + failed = [] + repaired = [] + quarantined = [] + count_by_group = Counter() + count_by_state = Counter() + cross = defaultdict(Counter) + + for i, sample in enumerate(samples, start=1): + try: + row = process_sample( + sample, + target_wh=(target_w, target_h), + classes=classes, + states=states, + nav_id=nav_id, + non_nav_id=non_nav_id, + dataset_root=DATASET_ROOT, + output_root=output_root, + repair_enabled=repair_enabled, + repair_max_unknown_pixels=int(args.repair_max_unknown_pixels), + repair_max_unknown_fraction=float(args.repair_max_unknown_fraction), + repair_color_max_distance=float(args.repair_color_max_distance), + repair_color_min_margin=float(args.repair_color_min_margin), + ) + processed.append(row) + + if row.get("mask_repair", {}).get("repaired"): + repaired.append({ + "group": sample.group, + "base": sample.base, + **row["mask_repair"], + }) + + count_by_group[row["group"]] += 1 + count_by_state[row["estado_corredor"]] += 1 + cross[row["group"]][row["estado_corredor"]] += 1 + + repair_tag = "" + + if row.get("mask_repair", {}).get("repaired"): + mr = row["mask_repair"] + repair_tag = ( + f" | REPAIR {mr['unknown_initial']}px " + f"(local={mr['repaired_spatial']}, " + f"color={mr['repaired_color']}, " + f"nearest={mr.get('repaired_nearest', 0)}, " + f"tieNonNav={mr.get('nearest_ties_to_non_nav', 0)})" + ) + + print( + f"[{i:5d}/{len(samples):5d}] OK " + f"{sample.group}/{sample.base} | " + f"mask={row['source_mask_contract']} -> ids | " + f"nav={row['navegavel_pct']:.1f}% | " + f"label={row['estado_corredor']}" + f"{repair_tag}" + ) + + except Exception as exc: + failure = { + "group": sample.group, + "base": sample.base, + "error_type": type(exc).__name__, + "error": str(exc), + } + failed.append(failure) + + quarantine_info = quarantine_sample( + sample, + quarantine_root=quarantine_root, + classes=classes, + error=exc, + ) + quarantined.append(quarantine_info) + + print( + f"[{i:5d}/{len(samples):5d}] QUARENTENA " + f"{sample.group}/{sample.base}: {exc}" + ) + + if fail_fast_processing: + raise + + report = { + "schema": REPORT_SCHEMA, + "created_at": now_iso(), + "config_path": str(config_path), + "labelmap_path": str(Path(args.labelmap)), + "input_root": str(original_group_root), + "output_root": str(output_root), + "target_resolution_wh": [target_w, target_h], + "classes": { + str(cid): { + "name": classes[cid][0], + "rgb": list(classes[cid][1]), + } + for cid in sorted(classes) + }, + "nav_class_id": nav_id, + "non_nav_class_id": non_nav_id, + "states": list(states), + "collection": collect_report, + "processed_count": len(processed), + "repaired_count": len(repaired), + "repaired": repaired, + "failed_count": len(failed), + "failed": failed, + "quarantined_count": len(quarantined), + "quarantine_root": str(quarantine_root), + "quarantined": quarantined, + "counts": { + "by_group": dict(count_by_group), + "by_state": dict(count_by_state), + "group_x_state": { + group: dict(counter) + for group, counter in cross.items() + }, + }, + "policy": { + "image_resize": "cv2.INTER_AREA", + "mask_input": "rgb_color_exact_or_legacy_ids", + "mask_repair_enabled": repair_enabled, + "mask_repair_policy": ( + "exact -> local_consensus -> strong_color_margin -> " + "nearest_valid_class(tie=non_nav) -> quarantine_if_over_limits" + ), + "repair_max_unknown_pixels": int(args.repair_max_unknown_pixels), + "repair_max_unknown_fraction": float(args.repair_max_unknown_fraction), + "repair_color_max_distance": float(args.repair_color_max_distance), + "repair_color_min_margin": float(args.repair_color_min_margin), + "mask_output": "uint8_class_ids", + "mask_resize": "cv2.INTER_NEAREST", + "preserve_groups": True, + "clear_downstream_on_clean": True, + "downstream_removed": removed_downstream, + "pairing": "canonical_legacy_stem + complete_triplets_only", + "incomplete_policy": "skip_and_report", + "collision_policy": "skip_and_report", + "complete_triplet_processing_error_policy": ( + "fail_fast" + if fail_fast_processing + else "quarantine_and_continue" + ), + "validate_folder_group_against_mask": True, + "augmentation_offline": False, + "mean_std_computed_here": False, + }, + } + + atomic_write_json(output_root / "normalize_report.json", report) + + print() + print("=" * 82) + print("NORMALIZE CONCLUÍDO") + print("=" * 82) + print(f"Processadas : {len(processed)}") + print(f"Reparadas mask : {len(repaired)}") + print(f"Quarentena : {len(quarantined)}") + print(f"Falhas conteúdo : {len(failed)}") + print(f"Incompletas skip : {collect_report['incomplete_count']}") + print(f"Aliases reconciliados: {collect_report['legacy_alias_count']}") + print(f"Colisões skip : {collect_report['collision_count']}") + print(f"Por grupo : {dict(count_by_group)}") + print(f"Por estado : {dict(count_by_state)}") + print(f"Saída : {output_root.resolve()}") + print(f"Revisão : {quarantine_root.resolve()}") + print("Próximo : split estratificado por grupo + estado, depois norm_stats do TRAIN") + print("=" * 82) + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/oak-d/_3_split_corridor.py b/Python/OAK/datasets/oak-d/_3_split_corridor.py new file mode 100644 index 000000000..d09bc1e3a --- /dev/null +++ b/Python/OAK/datasets/oak-d/_3_split_corridor.py @@ -0,0 +1,1682 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +_3_split_corridor.py +==================== + +Split inteligente do dataset normalizado da OAK-D Lite. + +Executar DE DENTRO da pasta oak-d: + + python _3_split_corridor.py + +Entrada: + config.json + + dataset/x/group/ + navegavel/ + images/ + masks/ + labels/ + naonavegavel/ + images/ + masks/ + labels/ + naonavegavel_navegavel/ + images/ + masks/ + labels/ + +Saída: + dataset/x/split/ + train/ + images/ + masks/ + labels/ + val/ + images/ + masks/ + labels/ + + dataset/x/ + norm_stats.json + split_report.json + +Objetivos: + 1) preservar a distribuição conjunta: + grupo da máscara x estado global do corredor + + 2) reduzir vazamento temporal: + uma sessão inteira de captura vai para train OU para val + + 3) calcular mean/std SOMENTE sobre train/images + + 4) não alterar dataset/x/group, que continua sendo a fonte + normalizada completa. + +Importante: + - o split é determinístico para a mesma seed; + - não usa sklearn; + - classes extremamente raras podem não conseguir aparecer nos dois lados; + isso é relatado explicitamente; + - se houver poucas sessões, a proporção train/val pode se afastar de 80/20; + isso é preferível a vazar frames quase idênticos entre os splits; + - labels copiados para split recebem metadados split/split_seed. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import os +import random +import re +import shutil +from collections import Counter, defaultdict +from dataclasses import dataclass, field +from datetime import datetime +from pathlib import Path +from typing import Dict, Iterable, List, Optional, Sequence, Tuple + +import cv2 +import numpy as np + + +# ============================================================================= +# Defaults +# ============================================================================= + +CONFIG_PATH = Path("config.json") +DATASET_ROOT = Path("dataset") + +DEFAULT_VAL_RATIO = 0.20 +DEFAULT_SEED = 42 + +# Blocos temporais: +# - abre bloco novo se a captura ficou muito tempo sem imagem; +# - mesmo numa sequência contínua, limita o span do bloco para evitar um +# bloco gigantesco que inviabilize o balanceamento. +DEFAULT_SESSION_GAP_S = 8.0 +DEFAULT_FALLBACK_BLOCK_SIZE = 12 + +REPORT_SCHEMA = "agrobot_corridor_split_report_v1" +NORM_SCHEMA = "agrobot_rgb_norm_stats_v2" + +IMG_EXTS = {".png", ".jpg", ".jpeg", ".bmp", ".webp"} +MASK_EXTS = {".png"} +LABEL_EXTS = {".json"} + + +# ============================================================================= +# Data +# ============================================================================= + +@dataclass(frozen=True) +class Sample: + base: str + group: str + image: Path + mask: Path + label: Path + label_id: int + label_name: str + timestamp_s: Optional[float] + capture_session_id: Optional[str] + + @property + def stratum(self) -> Tuple[str, int]: + return self.group, self.label_id + + +@dataclass +class Block: + block_id: int + samples: List[Sample] = field(default_factory=list) + + @property + def size(self) -> int: + return len(self.samples) + + @property + def strata(self) -> Counter: + return Counter(s.stratum for s in self.samples) + + @property + def first_timestamp(self) -> Optional[float]: + vals = [s.timestamp_s for s in self.samples if s.timestamp_s is not None] + return min(vals) if vals else None + + @property + def last_timestamp(self) -> Optional[float]: + vals = [s.timestamp_s for s in self.samples if s.timestamp_s is not None] + return max(vals) if vals else None + + +# ============================================================================= +# General helpers +# ============================================================================= + +def now_iso() -> str: + return datetime.now().isoformat(timespec="seconds") + + +def natural_key(text: str): + parts = re.split(r"(\d+)", str(text)) + return [int(p) if p.isdigit() else p.lower() for p in parts] + + +def safe_rel(path: Path, root: Path) -> str: + try: + return str(path.resolve().relative_to(root.resolve())).replace("\\", "/") + except Exception: + return str(path).replace("\\", "/") + + +def atomic_write_json(path: Path, data: dict) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_suffix(path.suffix + ".tmp") + + with tmp.open("w", encoding="utf-8") as f: + json.dump(data, f, ensure_ascii=False, indent=2) + + os.replace(tmp, path) + + +def atomic_write_csv(path: Path, rows: Sequence[dict], fieldnames: Sequence[str]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_suffix(path.suffix + ".tmp") + + with tmp.open("w", newline="", encoding="utf-8-sig") as f: + writer = csv.DictWriter( + f, + fieldnames=list(fieldnames), + extrasaction="ignore", + ) + writer.writeheader() + writer.writerows(rows) + + os.replace(tmp, path) + + +def map_by_stem(folder: Path, exts: set[str]) -> Dict[str, Path]: + if not folder.is_dir(): + return {} + + out: Dict[str, Path] = {} + + for p in sorted(folder.iterdir(), key=lambda x: natural_key(x.name)): + if not p.is_file() or p.suffix.lower() not in exts: + continue + + if p.stem in out: + raise RuntimeError( + f"Stem duplicado '{p.stem}' em {folder}: " + f"{out[p.stem].name} e {p.name}" + ) + + out[p.stem] = p + + return out + + +def parse_capture_session_id(base: str) -> Optional[str]: + """ + Novo padrão do capture: + img_sYYYYMMDD_HHMMSS__YYYYMMDD_HHMMSS_micro + + Retorna a identidade da execução do capture. + + Arquivos antigos não possuem essa informação e retornam None. + """ + m = re.search( + r"(?:^|_)s(\d{8}_\d{6})__", + str(base), + ) + if not m: + return None + return m.group(1) + + +def parse_capture_timestamp(base: str) -> Optional[float]: + """ + Suporta: + img_YYYYMMDD_HHMMSS_microseconds # legado + img_sSESSION__YYYYMMDD_HHMMSS_microseconds # atual + + Retorna POSIX timestamp local apenas para ordenar/calcular diferenças. + """ + matches = list( + re.finditer( + r"(\d{8})_(\d{6})(?:_(\d{1,6}))?", + str(base), + ) + ) + + if not matches: + return None + + # No formato novo, a primeira data é a sessão e a última é o frame. + m = matches[-1] + + date_txt = m.group(1) + time_txt = m.group(2) + frac_txt = m.group(3) or "0" + + frac_txt = (frac_txt + "000000")[:6] + + try: + dt = datetime.strptime( + date_txt + time_txt + frac_txt, + "%Y%m%d%H%M%S%f", + ) + return dt.timestamp() + except Exception: + return None + + +# ============================================================================= +# Read source dataset +# ============================================================================= + +def read_label(path: Path, states: Sequence[str]) -> Tuple[int, str, dict]: + try: + with path.open("r", encoding="utf-8") as f: + data = json.load(f) + except Exception as exc: + raise ValueError(f"Label JSON inválido {path}: {exc}") from exc + + if not isinstance(data, dict): + raise ValueError(f"Label não é objeto JSON: {path}") + + if "label_id" not in data: + raise ValueError(f"Label sem label_id: {path}") + + label_id = int(data["label_id"]) + + if not (0 <= label_id < len(states)): + raise ValueError( + f"label_id={label_id} fora de [0,{len(states)-1}] em {path}" + ) + + expected = str(states[label_id]) + got = data.get("estado_corredor") + + if got is not None and str(got) != expected: + raise ValueError( + f"estado_corredor inconsistente em {path}: " + f"id={label_id} => {expected!r}, arquivo={got!r}" + ) + + return label_id, expected, data + + +def collect_samples( + group_root: Path, + states: Sequence[str], +) -> List[Sample]: + if not group_root.is_dir(): + raise FileNotFoundError(f"group root não encontrado: {group_root}") + + result: List[Sample] = [] + seen_bases: Dict[str, str] = {} + + group_dirs = sorted( + [p for p in group_root.iterdir() if p.is_dir()], + key=lambda p: natural_key(p.name), + ) + + for group_dir in group_dirs: + group = group_dir.name + + images = map_by_stem(group_dir / "images", IMG_EXTS) + masks = map_by_stem(group_dir / "masks", MASK_EXTS) + labels = map_by_stem(group_dir / "labels", LABEL_EXTS) + + all_bases = set(images) | set(masks) | set(labels) + + incomplete = [] + + for base in sorted(all_bases, key=natural_key): + missing = [] + + if base not in images: + missing.append("image") + if base not in masks: + missing.append("mask") + if base not in labels: + missing.append("label") + + if missing: + incomplete.append((base, missing)) + continue + + if base in seen_bases: + raise RuntimeError( + f"Base duplicada entre grupos: {base!r} aparece em " + f"{seen_bases[base]!r} e {group!r}. " + "O split é plano e exige stems globalmente únicos." + ) + + seen_bases[base] = group + + label_id, label_name, _data = read_label( + labels[base], + states, + ) + + result.append( + Sample( + base=base, + group=group, + image=images[base], + mask=masks[base], + label=labels[base], + label_id=label_id, + label_name=label_name, + timestamp_s=parse_capture_timestamp(base), + capture_session_id=parse_capture_session_id(base), + ) + ) + + if incomplete: + preview = "\n".join( + f" {group}/{base}: faltando {','.join(missing)}" + for base, missing in incomplete[:20] + ) + + raise RuntimeError( + f"Grupo {group!r} possui {len(incomplete)} amostra(s) incompleta(s):\n" + + preview + ) + + if not result: + raise RuntimeError(f"Nenhuma amostra válida encontrada em {group_root}") + + return result + + +# ============================================================================= +# Temporal blocks +# ============================================================================= + +def build_capture_sessions( + samples: Sequence[Sample], + session_gap_s: float, + fallback_block_size: int, +) -> Tuple[List[Block], dict]: + """ + Constrói unidades indivisíveis de split. + + Prioridade: + 1) arquivos novos: usa capture_session_id explícito no nome; + 2) arquivos legados com timestamp: infere sessão por gap temporal; + 3) arquivos sem timestamp: fallback em blocos naturais. + + Regra central: + uma sessão inteira vai para TRAIN ou VAL. + + Isso é mais importante do que acertar exatamente 80/20. + Se existirem poucas sessões, o script aceita uma proporção menos perfeita + em vez de permitir vazamento temporal. + """ + explicit: Dict[str, List[Sample]] = defaultdict(list) + legacy_with_ts: List[Sample] = [] + without_ts: List[Sample] = [] + + for sample in samples: + if sample.capture_session_id: + explicit[str(sample.capture_session_id)].append(sample) + elif sample.timestamp_s is not None: + legacy_with_ts.append(sample) + else: + without_ts.append(sample) + + blocks: List[Block] = [] + next_id = 0 + + # --------------------------------------------------------- + # 1) Sessões explícitas do capture novo + # --------------------------------------------------------- + explicit_session_count = 0 + + for session_id in sorted(explicit, key=natural_key): + ss = sorted( + explicit[session_id], + key=lambda s: ( + float("inf") if s.timestamp_s is None else float(s.timestamp_s), + natural_key(s.base), + ), + ) + + blocks.append( + Block( + block_id=next_id, + samples=ss, + ) + ) + next_id += 1 + explicit_session_count += 1 + + # --------------------------------------------------------- + # 2) Dataset legado: sessão inferida por gap + # --------------------------------------------------------- + inferred_session_count = 0 + + if legacy_with_ts: + ordered = sorted( + legacy_with_ts, + key=lambda s: (float(s.timestamp_s), natural_key(s.base)), + ) + + current: List[Sample] = [] + prev_ts: Optional[float] = None + + for sample in ordered: + ts = float(sample.timestamp_s) + + if ( + current + and prev_ts is not None + and ts - prev_ts > session_gap_s + ): + blocks.append( + Block( + block_id=next_id, + samples=current, + ) + ) + next_id += 1 + inferred_session_count += 1 + current = [] + + current.append(sample) + prev_ts = ts + + if current: + blocks.append( + Block( + block_id=next_id, + samples=current, + ) + ) + next_id += 1 + inferred_session_count += 1 + + # --------------------------------------------------------- + # 3) Sem timestamp: fallback por blocos de ordenação natural + # --------------------------------------------------------- + fallback_block_count = 0 + + if without_ts: + ordered = sorted( + without_ts, + key=lambda s: natural_key(s.base), + ) + + bs = max(1, int(fallback_block_size)) + + for start_i in range(0, len(ordered), bs): + chunk = ordered[start_i:start_i + bs] + + blocks.append( + Block( + block_id=next_id, + samples=list(chunk), + ) + ) + next_id += 1 + fallback_block_count += 1 + + # Ordem cronológica aproximada só para legibilidade. + blocks.sort( + key=lambda b: ( + float("inf") if b.first_timestamp is None else b.first_timestamp, + b.block_id, + ) + ) + + blocks = [ + Block( + block_id=i, + samples=b.samples, + ) + for i, b in enumerate(blocks) + ] + + sizes = [b.size for b in blocks] + + report = { + "policy": "whole_capture_session_per_split", + "explicit_session_count": explicit_session_count, + "legacy_inferred_session_count": inferred_session_count, + "fallback_block_count": fallback_block_count, + "session_count_total": len(blocks), + + "explicit_session_samples": sum(len(v) for v in explicit.values()), + "legacy_timestamp_samples": len(legacy_with_ts), + "timestamp_missing_samples": len(without_ts), + + "session_size_min": min(sizes) if sizes else 0, + "session_size_max": max(sizes) if sizes else 0, + "session_size_mean": float(np.mean(sizes)) if sizes else 0.0, + + "legacy_session_gap_s": float(session_gap_s), + "fallback_block_size": int(fallback_block_size), + } + + return blocks, report + + +# ============================================================================= +# Intelligent split +# ============================================================================= + +def stratum_key(group: str, label_id: int) -> str: + return f"{group}__label_{label_id}" + + +def target_val_counts( + samples: Sequence[Sample], + val_ratio: float, +) -> Tuple[int, Counter, Counter]: + total = len(samples) + + target_total = int(round(total * val_ratio)) + + if total >= 2: + target_total = max(1, min(total - 1, target_total)) + else: + target_total = 0 + + total_strata = Counter(s.stratum for s in samples) + + target_strata: Counter = Counter() + + for stratum, count in total_strata.items(): + raw = count * val_ratio + + if count <= 1: + target = 0 + else: + target = int(round(raw)) + + # Sempre que possível, tenta manter ao menos 1 de cada lado. + target = max(1, target) + target = min(count - 1, target) + + target_strata[stratum] = target + + return target_total, total_strata, target_strata + + +def split_loss( + val_count: int, + val_strata: Counter, + target_total: int, + target_strata: Counter, + total_strata: Counter, +) -> float: + """ + Função de custo. + + - erro por stratum recebe peso maior quando a classe é rara; + - erro de tamanho total ajuda a segurar a razão global; + - overfill é permitido, mas custa. + """ + if target_total <= 0: + total_term = float(val_count > 0) + else: + total_term = abs(val_count - target_total) / max(1.0, float(target_total)) + + strata_term = 0.0 + weight_sum = 0.0 + + for st, total_n in total_strata.items(): + target_n = int(target_strata.get(st, 0)) + current_n = int(val_strata.get(st, 0)) + + denom = max(1.0, float(max(target_n, 1))) + err = abs(current_n - target_n) / denom + + rarity_weight = 1.0 / math.sqrt(max(1.0, float(total_n))) + # Multiplica por sqrt(N_total) para não deixar todos os pesos minúsculos. + rarity_weight *= math.sqrt(max(1.0, sum(total_strata.values()))) + + strata_term += err * rarity_weight + weight_sum += rarity_weight + + if weight_sum > 0: + strata_term /= weight_sum + + return ( + 0.72 * strata_term + + 0.28 * total_term + ) + + +def choose_val_blocks( + blocks: Sequence[Block], + samples: Sequence[Sample], + val_ratio: float, + seed: int, +) -> Tuple[set[int], dict]: + """ + Seleciona sessões inteiras para VAL. + + Prioridades, nesta ordem: + 1) quando um stratum (grupo × label) aparece em >=2 sessões, + tentar manter pelo menos uma ocorrência em TRAIN e VAL; + 2) aproximar a distribuição conjunta grupo × label; + 3) aproximar a razão global de validação. + + A proporção 80/20 é um ALVO, não uma regra mais importante que a + independência temporal ou a cobertura de casos. + """ + rng = random.Random(seed) + + target_total, total_strata, target_strata = target_val_counts( + samples, + val_ratio, + ) + + if target_total <= 0: + return set(), { + "target_total": target_total, + "target_strata": {}, + "coverage_feasible_strata": [], + "coverage_missing": [], + "loss": 0.0, + } + + # Em quantas sessões independentes cada stratum aparece? + stratum_blocks: Dict[Tuple[str, int], set[int]] = defaultdict(set) + + for block in blocks: + for st in block.strata: + stratum_blocks[st].add(block.block_id) + + feasible_both = { + st + for st, block_ids in stratum_blocks.items() + if len(block_ids) >= 2 + and total_strata.get(st, 0) >= 2 + } + + def coverage_missing( + val_strata: Counter, + ) -> List[Tuple[Tuple[str, int], str]]: + missing = [] + + for st in sorted(feasible_both): + val_n = int(val_strata.get(st, 0)) + total_n = int(total_strata.get(st, 0)) + + if val_n <= 0: + missing.append((st, "val")) + elif val_n >= total_n: + missing.append((st, "train")) + + return missing + + def objective( + val_count: int, + val_strata: Counter, + ) -> Tuple[int, float, int]: + missing = coverage_missing(val_strata) + + base_loss = split_loss( + val_count, + val_strata, + target_total, + target_strata, + total_strata, + ) + + size_error = abs( + int(val_count) + - int(target_total) + ) + + # Lexicográfico: + # cobertura primeiro, depois qualidade da estratificação, + # depois proximidade do tamanho. + return ( + len(missing), + float(base_loss), + int(size_error), + ) + + def calc_for_selected( + selected: set[int], + ) -> Tuple[int, Counter, Tuple[int, float, int]]: + vc = 0 + vs = Counter() + + for block in blocks: + if block.block_id in selected: + vc += block.size + vs += block.strata + + return vc, vs, objective(vc, vs) + + # ------------------------------------------------------------- + # Forward greedy + # ------------------------------------------------------------- + + remaining = list(blocks) + rng.shuffle(remaining) + + selected: set[int] = set() + val_count = 0 + val_strata: Counter = Counter() + current_obj = objective( + val_count, + val_strata, + ) + + while remaining: + candidates = [] + + for block in remaining: + new_count = val_count + block.size + new_strata = val_strata + block.strata + obj = objective( + new_count, + new_strata, + ) + + # O tamanho do bloco entra apenas como desempate final. + candidates.append( + ( + obj, + block.size, + block.block_id, + block, + new_count, + new_strata, + ) + ) + + candidates.sort( + key=lambda x: ( + x[0][0], + x[0][1], + x[0][2], + x[1], + x[2], + ) + ) + + ( + best_obj, + _best_size, + _best_id, + best_block, + new_count, + new_strata, + ) = candidates[0] + + # Continua enquanto: + # - ainda faltam casos cobríveis, OU + # - não atingimos o alvo total, OU + # - a inclusão melhora o objetivo. + need_coverage = current_obj[0] > 0 + need_size = val_count < target_total + improves = best_obj < current_obj + + if ( + val_count > 0 + and not need_coverage + and not need_size + and not improves + ): + break + + # Nunca joga todas as sessões em val. + if len(selected) + 1 >= len(blocks): + break + + selected.add( + best_block.block_id + ) + remaining = [ + b + for b in remaining + if b.block_id + != best_block.block_id + ] + + val_count = new_count + val_strata = new_strata + current_obj = best_obj + + if ( + val_count >= target_total + and current_obj[0] == 0 + ): + # Já temos cobertura e tamanho. + # Deixa o hill-climb encontrar ajustes melhores. + break + + # ------------------------------------------------------------- + # Hill-climb add/remove/swap + # ------------------------------------------------------------- + + if not selected: + best = min( + blocks, + key=lambda b: abs( + b.size - target_total + ), + ) + selected = { + best.block_id + } + + val_count, val_strata, current_obj = calc_for_selected( + selected + ) + + for _ in range(12): + best_sel = set(selected) + best_obj = current_obj + + selected_ids = sorted(selected) + unselected_ids = sorted( + b.block_id + for b in blocks + if b.block_id not in selected + ) + + # Remove. + for rid in selected_ids: + candidate = set(selected) + candidate.remove(rid) + + if not candidate: + continue + + _vc, _vs, obj = calc_for_selected( + candidate + ) + + if obj < best_obj: + best_obj = obj + best_sel = candidate + + # Add. + for aid in unselected_ids: + candidate = set(selected) + candidate.add(aid) + + if len(candidate) >= len(blocks): + continue + + _vc, _vs, obj = calc_for_selected( + candidate + ) + + if obj < best_obj: + best_obj = obj + best_sel = candidate + + # Swap. + # Para datasets enormes limita a vizinhança de busca. + max_side = 100 + + for rid in selected_ids[:max_side]: + for aid in unselected_ids[:max_side]: + candidate = set(selected) + candidate.remove(rid) + candidate.add(aid) + + if ( + not candidate + or len(candidate) >= len(blocks) + ): + continue + + _vc, _vs, obj = calc_for_selected( + candidate + ) + + if obj < best_obj: + best_obj = obj + best_sel = candidate + + if best_sel == selected: + break + + selected = best_sel + val_count, val_strata, current_obj = calc_for_selected( + selected + ) + + # Garante os dois lados mesmo em situação patológica. + if len(selected) >= len(blocks): + smallest = min( + ( + b + for b in blocks + if b.block_id in selected + ), + key=lambda b: b.size, + ) + selected.remove( + smallest.block_id + ) + + val_count, val_strata, current_obj = calc_for_selected( + selected + ) + + missing_final = coverage_missing( + val_strata + ) + + info = { + "target_total": int(target_total), + "actual_total": int(val_count), + + "target_strata": { + stratum_key(g, lid): int(v) + for (g, lid), v in sorted( + target_strata.items() + ) + }, + + "actual_strata": { + stratum_key(g, lid): int(v) + for (g, lid), v in sorted( + val_strata.items() + ) + }, + + "coverage_feasible_strata": [ + stratum_key(g, lid) + for (g, lid) in sorted( + feasible_both + ) + ], + + "coverage_missing": [ + { + "stratum": stratum_key( + st[0], + st[1], + ), + "missing_side": side, + } + for st, side in missing_final + ], + + "objective": { + "coverage_missing_count": int( + current_obj[0] + ), + "distribution_loss": float( + current_obj[1] + ), + "size_error": int( + current_obj[2] + ), + }, + } + + return selected, info + + +# ============================================================================= +# Distribution audit +# ============================================================================= + +def distribution(samples: Sequence[Sample], states: Sequence[str]) -> dict: + by_group = Counter(s.group for s in samples) + by_label = Counter(s.label_id for s in samples) + by_joint = Counter(s.stratum for s in samples) + + total = len(samples) + + return { + "total": total, + + "by_group": { + group: { + "count": int(count), + "pct": float(count * 100.0 / max(1, total)), + } + for group, count in sorted(by_group.items()) + }, + + "by_label": { + str(label_id): { + "name": str(states[label_id]), + "count": int(count), + "pct": float(count * 100.0 / max(1, total)), + } + for label_id, count in sorted(by_label.items()) + }, + + "by_group_label": { + stratum_key(group, label_id): { + "group": group, + "label_id": int(label_id), + "label_name": str(states[label_id]), + "count": int(count), + "pct": float(count * 100.0 / max(1, total)), + } + for (group, label_id), count in sorted(by_joint.items()) + }, + } + + +def rare_strata_warnings( + all_samples: Sequence[Sample], + train_samples: Sequence[Sample], + val_samples: Sequence[Sample], + states: Sequence[str], +) -> List[str]: + total = Counter(s.stratum for s in all_samples) + train = Counter(s.stratum for s in train_samples) + val = Counter(s.stratum for s in val_samples) + + warnings = [] + + for (group, label_id), count in sorted(total.items()): + t = int(train.get((group, label_id), 0)) + v = int(val.get((group, label_id), 0)) + + if count == 1: + warnings.append( + f"{group} × {states[label_id]} possui somente 1 amostra; " + f"ficou em {'train' if t else 'val'}." + ) + elif t == 0 or v == 0: + warnings.append( + f"{group} × {states[label_id]} possui {count} amostras, " + f"mas distribuição ficou train={t}, val={v}. " + "Blocos temporais podem impedir divisão sem vazamento." + ) + + return warnings + + +# ============================================================================= +# Copy split +# ============================================================================= + +def prepare_split_dirs(split_root: Path, clean: bool) -> Dict[str, Dict[str, Path]]: + if clean and split_root.exists(): + shutil.rmtree(split_root) + + result: Dict[str, Dict[str, Path]] = {} + + for split in ("train", "val"): + root = split_root / split + + result[split] = { + "root": root, + "images": root / "images", + "masks": root / "masks", + "labels": root / "labels", + } + + for key in ("images", "masks", "labels"): + result[split][key].mkdir(parents=True, exist_ok=True) + + return result + + +def copy_sample_to_split( + sample: Sample, + split: str, + dirs: Dict[str, Dict[str, Path]], + seed: int, + val_ratio: float, + source_root: Path, +) -> dict: + target = dirs[split] + + out_image = target["images"] / sample.image.name + out_mask = target["masks"] / f"{sample.base}.png" + out_label = target["labels"] / f"{sample.base}.json" + + shutil.copy2(sample.image, out_image) + shutil.copy2(sample.mask, out_mask) + + with sample.label.open("r", encoding="utf-8") as f: + label = json.load(f) + + label = dict(label) + label.update({ + "split": split, + "split_at": now_iso(), + "split_seed": int(seed), + "split_val_ratio": float(val_ratio), + "split_source_group": sample.group, + + "image": f"images/{out_image.name}", + "mask": f"masks/{out_mask.name}", + "label": f"labels/{out_label.name}", + + "source_normalized_image": safe_rel(sample.image, source_root), + "source_normalized_mask": safe_rel(sample.mask, source_root), + "source_normalized_label": safe_rel(sample.label, source_root), + }) + + atomic_write_json(out_label, label) + + return { + "base": sample.base, + "split": split, + "group": sample.group, + "label_id": sample.label_id, + "label_name": sample.label_name, + "timestamp_s": sample.timestamp_s, + "image": str(out_image), + "mask": str(out_mask), + "label": str(out_label), + } + + +# ============================================================================= +# Norm stats, TRAIN ONLY +# ============================================================================= + +def compute_train_norm_stats( + train_samples: Sequence[Sample], +) -> dict: + """ + Calcula estatísticas RGB em [0,1] sobre TODAS as imagens de train. + + Acumulação em float64 para evitar perda numérica. + """ + channel_sum = np.zeros(3, dtype=np.float64) + channel_sumsq = np.zeros(3, dtype=np.float64) + pixels = 0 + + shape_counts = Counter() + + for i, sample in enumerate(train_samples, start=1): + img_bgr = cv2.imread(str(sample.image), cv2.IMREAD_COLOR) + + if img_bgr is None: + raise FileNotFoundError( + f"Não consegui abrir imagem de train para norm_stats: {sample.image}" + ) + + img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB) + h, w, c = img_rgb.shape + + if c != 3: + raise RuntimeError( + f"Imagem não RGB 3 canais: {sample.image} shape={img_rgb.shape}" + ) + + shape_counts[(w, h)] += 1 + + arr = img_rgb.astype(np.float64) / 255.0 + + channel_sum += arr.sum(axis=(0, 1)) + channel_sumsq += np.square(arr).sum(axis=(0, 1)) + pixels += h * w + + if pixels <= 0: + raise RuntimeError("Train vazio; não é possível calcular norm_stats.") + + mean = channel_sum / pixels + var = channel_sumsq / pixels - np.square(mean) + std = np.sqrt(np.maximum(var, 1e-12)) + + return { + "schema": NORM_SCHEMA, + "created_at": now_iso(), + + "channels": ["R", "G", "B"], + "value_domain": "[0,1]", + + "mean": mean.tolist(), + "std": std.tolist(), + + "image_count": len(train_samples), + "pixels_per_channel": int(pixels), + + "source": "split/train/images_only", + "split_scope": "train_only", + + "image_shapes_wh": { + f"{w}x{h}": int(count) + for (w, h), count in sorted(shape_counts.items()) + }, + } + + +# ============================================================================= +# Temporal audit +# ============================================================================= + +def nearest_cross_split_time_gap( + train_samples: Sequence[Sample], + val_samples: Sequence[Sample], +) -> Optional[float]: + train_ts = sorted( + float(s.timestamp_s) + for s in train_samples + if s.timestamp_s is not None + ) + val_ts = sorted( + float(s.timestamp_s) + for s in val_samples + if s.timestamp_s is not None + ) + + if not train_ts or not val_ts: + return None + + i = 0 + j = 0 + best = float("inf") + + while i < len(train_ts) and j < len(val_ts): + best = min(best, abs(train_ts[i] - val_ts[j])) + + if train_ts[i] < val_ts[j]: + i += 1 + else: + j += 1 + + return float(best) + + +# ============================================================================= +# Main +# ============================================================================= + +def main(): + parser = argparse.ArgumentParser( + description=( + "Split inteligente grupo × label com proteção temporal " + "e norm_stats calculado somente no train." + ) + ) + + parser.add_argument( + "--config", + default=str(CONFIG_PATH), + ) + + parser.add_argument( + "--val-ratio", + type=float, + default=DEFAULT_VAL_RATIO, + ) + + parser.add_argument( + "--seed", + type=int, + default=DEFAULT_SEED, + ) + + parser.add_argument( + "--session-gap-s", + type=float, + default=DEFAULT_SESSION_GAP_S, + help=( + "Somente para arquivos legados sem session id: " + "gap que define nova sessão de captura." + ), + ) + + parser.add_argument( + "--fallback-block-size", + type=int, + default=DEFAULT_FALLBACK_BLOCK_SIZE, + help="Tamanho do bloco para nomes sem timestamp reconhecível.", + ) + + parser.add_argument( + "--no-clean", + action="store_true", + help="Não apaga split anterior antes de gerar.", + ) + + args = parser.parse_args() + + if not (0.0 < args.val_ratio < 0.5): + raise ValueError("--val-ratio deve ficar entre 0 e 0.5") + + if args.session_gap_s <= 0: + raise ValueError("--session-gap-s deve ser > 0") + + config_path = Path(args.config) + + if not config_path.is_file(): + raise FileNotFoundError(f"config não encontrado: {config_path}") + + with config_path.open("r", encoding="utf-8") as f: + config = json.load(f) + + resolution = config.get("resolucao") + + if ( + not isinstance(resolution, list) + or len(resolution) != 2 + ): + raise RuntimeError("config['resolucao'] deve ser [W,H]") + + W = int(resolution[0]) + H = int(resolution[1]) + + states = config.get("label_classes") + + if not isinstance(states, list) or not states: + raise RuntimeError("config['label_classes'] obrigatório.") + + normalized_root = DATASET_ROOT / f"{W}x{H}" + group_root = normalized_root / "group" + split_root = normalized_root / "split" + + samples = collect_samples( + group_root, + states, + ) + + blocks, temporal_report = build_capture_sessions( + samples, + session_gap_s=float(args.session_gap_s), + fallback_block_size=max(1, int(args.fallback_block_size)), + ) + + if len(blocks) < 2: + raise RuntimeError( + "O dataset possui menos de 2 sessões independentes de captura. " + "Não vou separar frames da mesma sessão entre train e val, pois isso " + "contaminaria a validação. Colete pelo menos mais uma sessão." + ) + + val_block_ids, split_algorithm = choose_val_blocks( + blocks, + samples, + val_ratio=float(args.val_ratio), + seed=int(args.seed), + ) + + train_samples: List[Sample] = [] + val_samples: List[Sample] = [] + + block_by_base = {} + + for block in blocks: + target = ( + val_samples + if block.block_id in val_block_ids + else train_samples + ) + + target.extend(block.samples) + + for sample in block.samples: + block_by_base[sample.base] = block.block_id + + # Ordem estável. + train_samples.sort(key=lambda s: natural_key(s.base)) + val_samples.sort(key=lambda s: natural_key(s.base)) + + if not train_samples or not val_samples: + raise RuntimeError( + f"Split inválido: train={len(train_samples)} val={len(val_samples)}" + ) + + dirs = prepare_split_dirs( + split_root, + clean=not args.no_clean, + ) + + print("=" * 88) + print("OAK-D Lite | Intelligent Corridor Split") + print("=" * 88) + print(f"Fonte normalizada : {group_root.resolve()}") + print(f"Saída split : {split_root.resolve()}") + print(f"Total : {len(samples)}") + print(f"Train : {len(train_samples)} ({len(train_samples)/len(samples)*100:.1f}%)") + print(f"Val : {len(val_samples)} ({len(val_samples)/len(samples)*100:.1f}%)") + print(f"Val alvo : {args.val_ratio*100:.1f}%") + print(f"Seed : {args.seed}") + print(f"Sessões de captura: {len(blocks)}") + print( + f"Política temporal : sessão inteira em apenas um split " + f"(legado gap>{args.session_gap_s:.1f}s)" + ) + print("=" * 88) + + train_set = {s.base for s in train_samples} + val_set = {s.base for s in val_samples} + + manifest_rows = [] + + for sample in samples: + split = "train" if sample.base in train_set else "val" + + row = copy_sample_to_split( + sample, + split=split, + dirs=dirs, + seed=int(args.seed), + val_ratio=float(args.val_ratio), + source_root=normalized_root, + ) + + row["capture_session_block"] = int(block_by_base[sample.base]) + row["capture_session_id"] = sample.capture_session_id or "" + manifest_rows.append(row) + + # Norm stats usa imagens de origem normalizada dos samples TRAIN. + norm_stats = compute_train_norm_stats( + train_samples + ) + + norm_stats.update({ + "resolution_wh": [W, H], + "split_seed": int(args.seed), + "split_val_ratio": float(args.val_ratio), + "train_count": len(train_samples), + "val_count": len(val_samples), + }) + + norm_path = normalized_root / "norm_stats.json" + atomic_write_json( + norm_path, + norm_stats, + ) + + dist_all = distribution(samples, states) + dist_train = distribution(train_samples, states) + dist_val = distribution(val_samples, states) + + warnings = rare_strata_warnings( + samples, + train_samples, + val_samples, + states, + ) + + nearest_gap = nearest_cross_split_time_gap( + train_samples, + val_samples, + ) + + if ( + nearest_gap is not None + and nearest_gap < args.session_gap_s + ): + warnings.append( + f"Menor distância temporal entre um frame train e val = " + f"{nearest_gap:.3f}s. Isso é menor que session_gap_s=" + f"{args.session_gap_s:.3f}s. Se os arquivos usam session id explícito, " + "isso significa que duas execuções distintas do capture ocorreram muito " + "próximas no tempo; revise se são realmente sessões independentes." + ) + + report = { + "schema": REPORT_SCHEMA, + "created_at": now_iso(), + + "config_path": str(config_path), + "normalized_root": str(normalized_root), + "group_root": str(group_root), + "split_root": str(split_root), + + "resolution_wh": [W, H], + + "seed": int(args.seed), + "val_ratio_requested": float(args.val_ratio), + "val_ratio_actual": float(len(val_samples) / len(samples)), + + "total_count": len(samples), + "train_count": len(train_samples), + "val_count": len(val_samples), + + "temporal_blocks": temporal_report, + "split_algorithm": split_algorithm, + + "distribution": { + "all": dist_all, + "train": dist_train, + "val": dist_val, + }, + + "nearest_train_val_timestamp_gap_s": nearest_gap, + + "warnings": warnings, + + "norm_stats": { + "path": str(norm_path), + "scope": "train_only", + "channels": norm_stats["channels"], + "mean": norm_stats["mean"], + "std": norm_stats["std"], + "pixels_per_channel": norm_stats["pixels_per_channel"], + }, + } + + report_path = normalized_root / "split_report.json" + + atomic_write_json( + report_path, + report, + ) + + manifest_path = normalized_root / "split_manifest.csv" + + atomic_write_csv( + manifest_path, + sorted(manifest_rows, key=lambda r: natural_key(r["base"])), + fieldnames=[ + "base", + "split", + "group", + "label_id", + "label_name", + "capture_session_block", + "capture_session_id", + "timestamp_s", + "image", + "mask", + "label", + ], + ) + + # ------------------------------------------------------------- + # Human-readable console audit + # ------------------------------------------------------------- + + print() + print("DISTRIBUIÇÃO POR GRUPO") + print("-" * 88) + + all_groups = sorted( + set(dist_all["by_group"]) + | set(dist_train["by_group"]) + | set(dist_val["by_group"]) + ) + + for group in all_groups: + a = dist_all["by_group"].get(group, {"count": 0})["count"] + t = dist_train["by_group"].get(group, {"count": 0})["count"] + v = dist_val["by_group"].get(group, {"count": 0})["count"] + + print( + f"{group:32s} total={a:5d} | train={t:5d} | val={v:5d}" + ) + + print() + print("DISTRIBUIÇÃO POR ESTADO") + print("-" * 88) + + for label_id, state in enumerate(states): + key = str(label_id) + + a = dist_all["by_label"].get(key, {"count": 0})["count"] + t = dist_train["by_label"].get(key, {"count": 0})["count"] + v = dist_val["by_label"].get(key, {"count": 0})["count"] + + print( + f"{label_id:2d} {str(state):28s} " + f"total={a:5d} | train={t:5d} | val={v:5d}" + ) + + print() + print("DISTRIBUIÇÃO CONJUNTA grupo × estado") + print("-" * 88) + + joint_keys = sorted(dist_all["by_group_label"]) + + for key in joint_keys: + info = dist_all["by_group_label"][key] + group = info["group"] + label_id = int(info["label_id"]) + state = info["label_name"] + + a = info["count"] + t = dist_train["by_group_label"].get(key, {"count": 0})["count"] + v = dist_val["by_group_label"].get(key, {"count": 0})["count"] + + print( + f"{group:28s} × {state:24s} " + f"total={a:4d} train={t:4d} val={v:4d}" + ) + + print() + print("NORM STATS - TRAIN ONLY") + print("-" * 88) + print(f"Arquivo : {norm_path.resolve()}") + print(f"Mean RGB: {[round(x, 8) for x in norm_stats['mean']]}") + print(f"Std RGB: {[round(x, 8) for x in norm_stats['std']]}") + + if nearest_gap is not None: + print(f"Menor gap temporal train↔val: {nearest_gap:.3f}s") + + if warnings: + print() + print("AVISOS") + print("-" * 88) + for warning in warnings: + print(f"[WARN] {warning}") + + print() + print("=" * 88) + print("SPLIT CONCLUÍDO") + print("=" * 88) + print(f"Train : {dirs['train']['root'].resolve()}") + print(f"Val : {dirs['val']['root'].resolve()}") + print(f"norm_stats : {norm_path.resolve()}") + print(f"split_report : {report_path.resolve()}") + print(f"split_manifest : {manifest_path.resolve()}") + print("=" * 88) + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/oak-d/_4_train_corridor.py b/Python/OAK/datasets/oak-d/_4_train_corridor.py new file mode 100644 index 000000000..22cf1d95e --- /dev/null +++ b/Python/OAK/datasets/oak-d/_4_train_corridor.py @@ -0,0 +1,4198 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +Agri Corridor Teacher v2 +======================== + +Trainer especializado para a camera frontal OAK-D Lite do robo agricola. + +Tarefa operacional: + 1) Segmentacao binaria: + - area nao navegavel + - area navegavel + + 2) Classificacao global do estado do corredor: + - classes declaradas em config["label_classes"] + +Filosofia: + - nao tenta ser um trainer generico; + - prioriza robustez em campo; + - augmentation preserva a semantica global do corredor; + - loss de segmentacao combina CE + Dice + borda + penalidade de falso navegavel; + - status head trabalha em features nativas e resumo espacial da segmentacao, + sem fazer upsample gigante de features profundas; + - sampler ajuda a balancear estados raros sem explodir repeticoes; + - LR discriminativo entre encoder, decoder e status head; + - warmup + polynomial decay; + - AMP, grad accumulation e gradient clipping; + - checkpoints separados para navegacao, status e score operacional; + - best_operational so nasce quando a epoca passa pelos gates de campo; + - consome SOMENTE o split final em dataset/x/split; + - norm_stats deve vir do split e ser train-only; + - sampler conjunto grupo-da-mascara x estado global; + - preflight audita paths, shapes, IDs, labels e vazamento de stems; + - metricas de validacao tambem sao quebradas por grupo da mascara; + - checkpoint salva o contrato necessario para exportacao ONNX posterior; + - persiste historico completo por epoca em JSONL + JSON + CSV. + +Dependencias: + numpy + pillow + torch + transformers + roi_seg_dataset.ROISegDataset + +Exemplo: + python train_corridor_agri_v2.py ^ + --config config.json ^ + --epochs 140 ^ + --batch 4 ^ + --num_workers 4 ^ + --amp ^ + --amp_val ^ + --grad_accum 2 +""" + +from __future__ import annotations + +import argparse +import copy +import csv +import json +import math +import os +import random +import time +from collections import Counter +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Dict, List, Optional, Sequence, Tuple + +import cv2 +import numpy as np +from PIL import Image + +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch.cuda.amp import GradScaler +from torch.utils.data import DataLoader, WeightedRandomSampler +from transformers import SegformerForSemanticSegmentation + +from roi_seg_dataset import ROISegDataset + + +# ============================================================================= +# Contrato / defaults +# ============================================================================= + +TRAINER_VERSION = "agri_corridor_teacher_v2.1_operational_gate" +IGNORE_INDEX = 255 +LABEL_IGNORE_INDEX = -100 + +DEFAULT_TRAINING = { + "data": { + # Campo: label de status ausente normalmente significa dataset inconsistente. + "strict_labels": True, + + # O pipeline novo exige split/ e norm_stats produzidos pela etapa _3_split. + "strict_pipeline_contract": True, + + # Quantas amostras de cada split auditar profundamente antes do treino. + # 0 = todas. + "preflight_max_samples": 0, + }, + + "augmentation": { + "enabled": True, + + # Geometria segura para corredor. + # Horizontal flip e pequenas perturbacoes preservam o significado global. + "horizontal_flip_p": 0.50, + "affine_p": 0.60, + "rotate_deg": 4.0, + "scale_min": 0.92, + "scale_max": 1.08, + "translate_frac": 0.03, + + # Fotometria pensada para sol, sombra e autoexposure no campo. + "global_gain_p": 0.45, + "global_gain_min": 0.82, + "global_gain_max": 1.18, + + "contrast_p": 0.35, + "contrast_min": 0.85, + "contrast_max": 1.15, + + "gamma_p": 0.30, + "gamma_min": 0.82, + "gamma_max": 1.18, + + "rgb_gain_p": 0.25, + "rgb_gain_min": 0.94, + "rgb_gain_max": 1.06, + + "shadow_p": 0.35, + "shadow_strength_min": 0.12, + "shadow_strength_max": 0.42, + + "noise_p": 0.18, + "noise_sigma_min": 0.003, + "noise_sigma_max": 0.018, + + "blur_p": 0.12, + "blur_kernel": 3, + + # Nao habilitado por padrao. Pode ser usado futuramente para lente suja/folha. + "occlusion_p": 0.00, + "occlusion_min_frac": 0.04, + "occlusion_max_frac": 0.15, + + # Comeca leve e chega a 100% da intensidade. + "ramp_epochs": 6, + }, + + "sampler": { + # O dataset possui dois eixos pedagogicos: + # 1) composicao da mascara: navegavel / nao navegavel / misto + # 2) estado global: EntrandoRua / CaminhandoRua / ... + # + # "joint" balanceia a combinacao grupo x estado. + # Se metadado de grupo faltar, cai automaticamente para status-only. + "mode": "joint", + "joint_power": 0.30, + "status_power": 0.12, + "group_power": 0.10, + "max_weight_ratio": 5.0, + "samples_per_epoch": 0, + }, + + "class_weighting": { + # Segmentacao. + "seg_method": "log_inverse", + "seg_log_offset": 1.02, + "seg_min_weight": 0.35, + "seg_max_weight": 3.0, + "seg_max_samples": 1200, + + # Status. + "status_enabled": True, + "status_power": 0.30, + "status_min_weight": 0.50, + "status_max_weight": 3.0, + }, + + "loss": { + # Segmentacao. + "seg_ce_weight": 0.65, + "seg_dice_weight": 0.35, + + # Aumenta a atencao na fronteira corredor / nao corredor. + "boundary_boost": 0.15, + + # Penaliza confianca navegavel sobre GT nao navegavel. + # Mantido pequeno para nao ensinar o modelo a "ter medo de tudo". + "unsafe_nav_weight": 0.06, + + # Status global. + "status_weight": 0.40, + "status_ramp_epochs": 8, + "status_label_smoothing": 0.03, + }, + + "status_head": { + # Resumo espacial compacto. Preserva geometria global sem upsample de feature. + "pool_h": 3, + "pool_w": 4, + "hidden": 256, + "dropout": 0.20, + + # Evita que a loss de status force o decoder semantico a "codificar status". + # A label head ainda ensina o encoder atraves de feat. + "detach_seg_summary": True, + + # Usa probabilidade, nao logits crus, como resumo da geometria semantica. + "seg_summary": "probabilities", + }, + + "optimizer": { + "encoder_lr_mult": 0.50, + "decoder_lr_mult": 1.00, + "status_lr_mult": 2.00, + "weight_decay": None, # None => usa --wd + "no_decay_bias": True, + "no_decay_norm": True, + "betas": [0.9, 0.999], + "eps": 1e-8, + + # Primeiro deixa heads se organizarem sem mexer no backbone inteiro. + "freeze_encoder_epochs": 2, + }, + + "scheduler": { + "mode": "poly", + "warmup_ratio": 0.05, + "warmup_start_factor": 0.10, + "poly_power": 1.0, + "min_lr_ratio": 0.02, + }, + + "optimization": { + "grad_clip_norm": 1.0, + "matmul_precision": "high", + "cudnn_benchmark": True, + "persistent_workers": True, + "prefetch_factor": 2, + }, + + "score": { + # O FIELD_SCORE RANQUEIA somente os modelos que passaram pelo gate. + "nav_iou": 0.40, + "nav_f1": 0.20, + "status_macro_f1": 0.20, + "safety": 0.20, + }, + + "operational_gate": { + # 1) elegibilidade para campo; 2) FIELD_SCORE escolhe o melhor elegivel. + "enabled": True, + "min_epoch": 5, + "min_nav_iou": 0.70, + "min_nav_f1": 0.80, + "min_non_nav_iou": 0.60, + "min_status_macro_f1": 0.60, + "max_unsafe_nav_rate": 0.03, + + "group_unsafe": { + "enabled": True, + "min_non_nav_pixels": 5000, + "max_rate": 0.05, + "max_rate_by_group": { + "naonavegavel": 0.04, + "navegavel_naonavegavel": 0.05 + } + }, + + "status_recall": { + "enabled": False, + "min_support": 5, + "minimums": {} + } + }, + + "checkpoint": { + "early_stop_patience": 18, + "early_stop_min_delta": 3e-4, + }, +} + + +# ============================================================================= +# Utilitarios +# ============================================================================= + +def deep_merge(base: dict, override: Optional[dict]) -> dict: + out = copy.deepcopy(base) + if not isinstance(override, dict): + return out + + def rec(dst: dict, src: dict): + for k, v in src.items(): + if isinstance(v, dict) and isinstance(dst.get(k), dict): + rec(dst[k], v) + else: + dst[k] = copy.deepcopy(v) + + rec(out, override) + return out + + +def seed_everything(seed: int = 42): + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + + +def seed_worker(worker_id: int): + seed = torch.initial_seed() % (2 ** 32) + np.random.seed(seed) + random.seed(seed) + + +def capture_rng_state() -> dict: + state = { + "python": random.getstate(), + "numpy": np.random.get_state(), + "torch": torch.get_rng_state(), + } + if torch.cuda.is_available(): + state["cuda"] = torch.cuda.get_rng_state_all() + return state + + +def restore_rng_state(state: Optional[dict]): + if not state: + return + if "python" in state: + random.setstate(state["python"]) + if "numpy" in state: + np.random.set_state(state["numpy"]) + if "torch" in state: + torch.set_rng_state(state["torch"]) + if torch.cuda.is_available() and "cuda" in state: + torch.cuda.set_rng_state_all(state["cuda"]) + + +def ensure_img_tensor(img: Any) -> torch.Tensor: + if isinstance(img, Image.Image): + img = np.asarray(img) + + if isinstance(img, np.ndarray): + img = torch.from_numpy(img) + + if not isinstance(img, torch.Tensor): + raise TypeError(f"Imagem com tipo inesperado: {type(img)}") + + if img.ndim == 3 and img.shape[-1] == 3: + img = img.permute(2, 0, 1) + + if img.ndim != 3 or img.shape[0] != 3: + raise RuntimeError(f"Imagem deve ser CHW RGB, recebido shape={tuple(img.shape)}") + + img = img.float() + if img.numel() > 0 and float(img.max()) > 1.5: + img = img / 255.0 + + return img.contiguous().clamp(0.0, 1.0) + + +def ensure_mask_tensor(mask: Any) -> torch.Tensor: + if isinstance(mask, Image.Image): + mask = np.asarray(mask) + + if isinstance(mask, np.ndarray): + mask = torch.from_numpy(mask) + + if not isinstance(mask, torch.Tensor): + raise TypeError(f"Mask com tipo inesperado: {type(mask)}") + + if mask.ndim == 3: + if mask.shape[0] == 1: + mask = mask[0] + elif mask.shape[-1] == 1: + mask = mask[..., 0] + else: + mask = mask[..., 0] + + return mask.long().contiguous() + + +def read_labelmap(path: Path) -> Tuple[Dict[int, str], Dict[str, int]]: + if not path.is_file(): + raise FileNotFoundError(f"labelmap nao encontrado: {path}") + + id2label: Dict[int, str] = {} + + with path.open("r", encoding="utf-8") as f: + for raw in f: + line = raw.strip() + if not line or line.startswith("#"): + continue + + parts = [p.strip() for p in line.replace("=", ":").split(":") if p.strip()] + + if len(parts) >= 2: + try: + cid = int(parts[0]) + name = parts[1] + except ValueError: + cid = len(id2label) + name = parts[0] + else: + cid = len(id2label) + name = parts[0] + + if name.lower() in ("ignore", "void", "background_ignore"): + continue + + id2label[int(cid)] = str(name) + + label2id = {name.lower(): cid for cid, name in id2label.items()} + return id2label, label2id + + +def resolve_project_root( + config_path: Path, + camera_value: Optional[str] = None, +) -> Path: + """ + O pipeline 2026-09 roda de DENTRO da pasta da camera (oak-d). + + Regra principal: + config.json.parent precisa conter dataset/. + + Fallback apenas para compatibilidade: + tenta camera_value se o config for chamado de outro local. + + Nao cria caminhos do tipo oak-d/oak-d. + """ + primary = config_path.parent.resolve() + + if (primary / "dataset").is_dir(): + return primary + + candidates: List[Path] = [] + + if camera_value: + raw = Path(str(camera_value)) + + if raw.is_absolute(): + candidates.append(raw.resolve()) + else: + candidates.extend([ + (Path.cwd() / raw).resolve(), + (config_path.parent / raw).resolve(), + ]) + + # cwd tambem pode ser a raiz, mesmo se config foi fornecido com caminho absoluto. + candidates.append(Path.cwd().resolve()) + + seen = set() + + for cand in candidates: + key = str(cand) + if key in seen: + continue + seen.add(key) + + if (cand / "dataset").is_dir(): + return cand + + pretty = "\n".join(f" - {p}" for p in [primary] + candidates) + + raise FileNotFoundError( + "Nao consegui resolver a raiz do projeto/camera. " + "O trainer novo deve ser executado dentro de oak-d, com dataset/ ao lado " + f"de config.json.\nCaminhos testados:\n{pretty}" + ) + + +def load_norm_stats( + path: Path, + device: torch.device, + expected_resolution_wh: Tuple[int, int], + strict_contract: bool = True, +) -> Tuple[torch.Tensor, torch.Tensor, Path, dict]: + """ + Carrega o norm_stats produzido pela etapa _3_split_corridor.py. + + Contrato novo: + dataset/x/norm_stats.json + split_scope == "train_only" + channels == ["R","G","B"] + resolution_wh == [W,H] + """ + path = Path(path) + + if not path.is_file(): + raise FileNotFoundError( + f"norm_stats.json obrigatorio nao encontrado: {path}\n" + "Execute _3_split_corridor.py antes do treino." + ) + + with path.open("r", encoding="utf-8") as f: + stats = json.load(f) + + channels = [str(x).upper() for x in stats.get("channels", [])] + mean = list(stats.get("mean", [])) + std = list(stats.get("std", [])) + + if channels != ["R", "G", "B"]: + raise ValueError( + f"norm_stats precisa ter channels=['R','G','B']; recebido={channels}" + ) + + if len(mean) != 3 or len(std) != 3: + raise ValueError(f"mean/std invalidos: mean={mean} std={std}") + + if any((not np.isfinite(float(x))) for x in mean + std): + raise ValueError("norm_stats contem NaN/Inf.") + + if any(float(x) <= 0 for x in std): + raise ValueError(f"std precisa ser > 0; recebido={std}") + + split_scope = str(stats.get("split_scope", "")).strip().lower() + + if strict_contract and split_scope != "train_only": + raise RuntimeError( + f"norm_stats fora do contrato novo: split_scope={split_scope!r}. " + "Esperado 'train_only'. Reexecute _3_split_corridor.py." + ) + + expected_wh = [int(expected_resolution_wh[0]), int(expected_resolution_wh[1])] + stats_wh = stats.get("resolution_wh") + + if strict_contract and stats_wh is not None: + got_wh = [int(stats_wh[0]), int(stats_wh[1])] + + if got_wh != expected_wh: + raise RuntimeError( + f"norm_stats resolution_wh={got_wh}, mas config.resolucao={expected_wh}." + ) + + value_domain = str(stats.get("value_domain", "[0,1]")).replace(" ", "") + + if strict_contract and value_domain not in {"[0,1]", "0..1", "0-1"}: + raise RuntimeError( + f"norm_stats value_domain inesperado: {stats.get('value_domain')!r}" + ) + + mean_t = torch.tensor( + mean, + dtype=torch.float32, + device=device, + ).view(1, 3, 1, 1) + + std_t = torch.tensor( + std, + dtype=torch.float32, + device=device, + ).view(1, 3, 1, 1).clamp_min(1e-6) + + return mean_t, std_t, path, stats + + +# ============================================================================= +# Dataset de corredor + status +# ============================================================================= + +class CorridorDataset(torch.utils.data.Dataset): + """ + Wrapper do ROISegDataset para o contrato final do split. + + Esperado: + split_root/ + images/.png + masks/.png + labels/.json + + O JSON do split contem: + label_id + estado_corredor + split_source_group + + Mantemos o grupo porque ele e um eixo pedagogico importante tanto para + sampler quanto para metricas de validacao. + """ + + def __init__( + self, + base_ds: ROISegDataset, + split_root: Path, + label_classes: Sequence[str], + strict_labels: bool = True, + ): + self.base_ds = base_ds + self.split_root = Path(split_root) + self.label_classes = [str(x) for x in label_classes] + self.num_label_classes = len(self.label_classes) + self.strict_labels = bool(strict_labels) + + self.label_name_to_id = { + str(name).strip().lower(): i + for i, name in enumerate(self.label_classes) + } + + self.label_ids: List[int] = [] + self.label_paths: List[Optional[str]] = [] + self.group_names: List[str] = [] + self.group_ids: List[int] = [] + self.missing_labels: List[str] = [] + + raw_groups: List[str] = [] + + for idx in range(len(self.base_ds)): + mask_path = self._get_mask_path(idx) + label_path = self._label_path_from_mask_path(mask_path) + + label_id = LABEL_IGNORE_INDEX + label_name: Optional[str] = None + group_name = "unknown" + + if label_path is None or not os.path.exists(label_path): + self.missing_labels.append(mask_path) + else: + label_id_opt, label_name, group_name = self._read_label(label_path) + + if label_id_opt is None and label_name: + label_id_opt = self.label_name_to_id.get( + str(label_name).strip().lower() + ) + + if label_id_opt is None: + self.missing_labels.append(mask_path) + else: + label_id = int(label_id_opt) + + if label_id < 0 or label_id >= self.num_label_classes: + raise ValueError( + f"label_id={label_id} fora do intervalo " + f"[0,{self.num_label_classes - 1}] em {label_path}" + ) + + expected_name = self.label_classes[label_id] + + if ( + label_name is not None + and str(label_name) != str(expected_name) + ): + raise ValueError( + f"Label inconsistente em {label_path}: " + f"label_id={label_id} implica {expected_name!r}, " + f"mas nome={label_name!r}." + ) + + self.label_ids.append(int(label_id)) + self.label_paths.append(label_path) + raw_groups.append(str(group_name or "unknown")) + + if self.strict_labels and self.missing_labels: + preview = "\n".join( + f" - {p}" + for p in self.missing_labels[:10] + ) + + raise RuntimeError( + f"{len(self.missing_labels)} amostras sem label de status " + f"no split {self.split_root}.\nPrimeiras:\n{preview}\n" + "Corrija o split/dataset ou use --allow_missing_labels." + ) + + unique_groups = sorted(set(raw_groups)) + self.group_name_by_id = { + i: name + for i, name in enumerate(unique_groups) + } + self.group_id_by_name = { + name: i + for i, name in self.group_name_by_id.items() + } + + self.group_names = raw_groups + self.group_ids = [ + self.group_id_by_name[g] + for g in raw_groups + ] + + def __len__(self): + return len(self.base_ds) + + def _get_mask_path(self, idx: int) -> str: + if hasattr(self.base_ds, "msk_paths"): + return str(self.base_ds.msk_paths[idx]) + + item = self.base_ds[idx] + + if isinstance(item, dict): + for key in ("mask_path", "mask_file", "maskname", "mask_name"): + if item.get(key): + return str(item[key]) + + raise RuntimeError("ROISegDataset nao expoe caminho da mask.") + + def _label_path_from_mask_path( + self, + mask_path: str, + ) -> Optional[str]: + norm = os.path.normpath(mask_path) + parts = norm.split(os.sep) + stem = os.path.splitext(os.path.basename(mask_path))[0] + + try: + i = parts.index("masks") + root = os.sep.join(parts[:i] + ["labels"]) + except ValueError: + root = str(self.split_root / "labels") + + # Pipeline novo usa JSON. NPY/TXT ficam apenas como compatibilidade. + for ext in (".json", ".npy", ".txt"): + candidate = os.path.join(root, stem + ext) + if os.path.exists(candidate): + return candidate + + return os.path.join(root, stem + ".json") + + @staticmethod + def _read_label( + path: str, + ) -> Tuple[Optional[int], Optional[str], str]: + ext = os.path.splitext(path)[1].lower() + + if ext == ".npy": + value = np.load(path) + return int(np.asarray(value).reshape(-1)[0]), None, "unknown" + + if ext == ".json": + with open(path, "r", encoding="utf-8") as f: + data = json.load(f) + + label_id = data.get("label_id") + + label_name = ( + data.get("estado_corredor") + or data.get("label") + or data.get("state") + ) + + group_name = ( + data.get("split_source_group") + or data.get("group") + or data.get("source_group") + or "unknown" + ) + + return ( + int(label_id) if label_id is not None else None, + str(label_name) if label_name is not None else None, + str(group_name), + ) + + if ext == ".txt": + txt = Path(path).read_text( + encoding="utf-8" + ).strip() + + try: + return int(txt), None, "unknown" + except ValueError: + return None, txt, "unknown" + + return None, None, "unknown" + + def __getitem__(self, idx: int): + item = self.base_ds[idx] + + if ( + isinstance(item, (tuple, list)) + and len(item) == 1 + and isinstance(item[0], dict) + ): + item = item[0] + + if isinstance(item, dict): + img = item["image"] + mask = item["mask"] + elif isinstance(item, (tuple, list)) and len(item) >= 2: + img, mask = item[:2] + else: + raise RuntimeError( + f"Formato inesperado de item: {type(item)}" + ) + + return { + "image": ensure_img_tensor(img), + "mask": ensure_mask_tensor(mask), + "label": int(self.label_ids[idx]), + "group_id": int(self.group_ids[idx]), + "group_name": self.group_names[idx], + "label_path": self.label_paths[idx], + "mask_path": self._get_mask_path(idx), + } + + +def collate_corridor(batch: Sequence[dict]): + imgs = torch.stack( + [x["image"] for x in batch], + dim=0, + ) + masks = torch.stack( + [x["mask"] for x in batch], + dim=0, + ) + labels = torch.tensor( + [int(x["label"]) for x in batch], + dtype=torch.long, + ) + group_ids = torch.tensor( + [int(x["group_id"]) for x in batch], + dtype=torch.long, + ) + + return imgs, masks, labels, group_ids + + +# ============================================================================= +# Augmentation de campo +# ============================================================================= + +class FieldCorridorAugmenter: + """ + Augmentation sincronizado imagem/mask. + + Importante: + - nao faz vertical flip; + - nao faz random crop por padrao; + - perturbacoes geometricas sao pequenas; + - status global permanece semanticamente valido. + """ + + def __init__(self, cfg: dict, ignore_index: int = IGNORE_INDEX): + self.cfg = cfg or {} + self.ignore_index = int(ignore_index) + self.strength = 1.0 + + def set_epoch(self, epoch: int): + ramp = int(self.cfg.get("ramp_epochs", 0)) + if ramp <= 0: + self.strength = 1.0 + else: + self.strength = min(1.0, max(0.20, float(epoch) / float(ramp))) + + def __call__( + self, + imgs: torch.Tensor, + masks: torch.Tensor, + ) -> Tuple[torch.Tensor, torch.Tensor]: + if not bool(self.cfg.get("enabled", True)): + return imgs, masks + + imgs = imgs.clone() + masks = masks.clone() + + imgs, masks = self._geometry(imgs, masks) + imgs = self._photometric(imgs) + + return imgs.clamp(0.0, 1.0), masks + + def _geometry( + self, + imgs: torch.Tensor, + masks: torch.Tensor, + ) -> Tuple[torch.Tensor, torch.Tensor]: + B, _, H, W = imgs.shape + device = imgs.device + + # Horizontal flip. + p_flip = float(self.cfg.get("horizontal_flip_p", 0.5)) * self.strength + do_flip = torch.rand(B, device=device) < p_flip + + if bool(do_flip.any()): + imgs[do_flip] = torch.flip(imgs[do_flip], dims=[3]) + masks[do_flip] = torch.flip(masks[do_flip], dims=[2]) + + # Affine leve. + p_aff = float(self.cfg.get("affine_p", 0.6)) * self.strength + do_aff = torch.rand(B, device=device) < p_aff + + if not bool(do_aff.any()): + return imgs, masks + + rotate_deg = float(self.cfg.get("rotate_deg", 4.0)) * self.strength + + smin = float(self.cfg.get("scale_min", 0.92)) + smax = float(self.cfg.get("scale_max", 1.08)) + smin = 1.0 - (1.0 - smin) * self.strength + smax = 1.0 + (smax - 1.0) * self.strength + + trans = float(self.cfg.get("translate_frac", 0.03)) * self.strength + + n = int(do_aff.sum().item()) + angle = (torch.rand(n, device=device) * 2.0 - 1.0) * math.radians(rotate_deg) + scale = torch.empty(n, device=device).uniform_(smin, smax) + tx = (torch.rand(n, device=device) * 2.0 - 1.0) * (2.0 * trans) + ty = (torch.rand(n, device=device) * 2.0 - 1.0) * (2.0 * trans) + + cos = torch.cos(angle) * scale + sin = torch.sin(angle) * scale + + theta = torch.zeros((n, 2, 3), dtype=imgs.dtype, device=device) + theta[:, 0, 0] = cos + theta[:, 0, 1] = -sin + theta[:, 1, 0] = sin + theta[:, 1, 1] = cos + theta[:, 0, 2] = tx + theta[:, 1, 2] = ty + + img_sel = imgs[do_aff] + mask_sel = masks[do_aff].float().unsqueeze(1) + + grid = F.affine_grid( + theta, + size=img_sel.size(), + align_corners=False, + ) + + img_aug = F.grid_sample( + img_sel, + grid, + mode="bilinear", + padding_mode="border", + align_corners=False, + ) + + mask_aug = F.grid_sample( + mask_sel, + grid, + mode="nearest", + padding_mode="zeros", + align_corners=False, + ).squeeze(1).long() + + # Marca pixels fora da imagem como ignore, em vez de ensinar classe 0 falsa. + valid_src = torch.ones( + (n, 1, H, W), + dtype=imgs.dtype, + device=device, + ) + valid_aug = F.grid_sample( + valid_src, + grid, + mode="nearest", + padding_mode="zeros", + align_corners=False, + ).squeeze(1) > 0.5 + + mask_aug[~valid_aug] = self.ignore_index + + imgs[do_aff] = img_aug + masks[do_aff] = mask_aug + + return imgs, masks + + def _photometric(self, imgs: torch.Tensor) -> torch.Tensor: + B, C, H, W = imgs.shape + device = imgs.device + dtype = imgs.dtype + + def choose(p: float) -> torch.Tensor: + return torch.rand(B, device=device) < (float(p) * self.strength) + + # Ganho global / exposicao. + sel = choose(self.cfg.get("global_gain_p", 0.45)) + if bool(sel.any()): + lo = float(self.cfg.get("global_gain_min", 0.82)) + hi = float(self.cfg.get("global_gain_max", 1.18)) + lo = 1.0 - (1.0 - lo) * self.strength + hi = 1.0 + (hi - 1.0) * self.strength + gain = torch.empty((int(sel.sum()), 1, 1, 1), device=device, dtype=dtype).uniform_(lo, hi) + imgs[sel] = imgs[sel] * gain + + # Contraste. + sel = choose(self.cfg.get("contrast_p", 0.35)) + if bool(sel.any()): + lo = float(self.cfg.get("contrast_min", 0.85)) + hi = float(self.cfg.get("contrast_max", 1.15)) + lo = 1.0 - (1.0 - lo) * self.strength + hi = 1.0 + (hi - 1.0) * self.strength + contrast = torch.empty((int(sel.sum()), 1, 1, 1), device=device, dtype=dtype).uniform_(lo, hi) + mean = imgs[sel].mean(dim=(2, 3), keepdim=True) + imgs[sel] = (imgs[sel] - mean) * contrast + mean + + # Gamma. + sel = choose(self.cfg.get("gamma_p", 0.30)) + if bool(sel.any()): + lo = float(self.cfg.get("gamma_min", 0.82)) + hi = float(self.cfg.get("gamma_max", 1.18)) + lo = 1.0 - (1.0 - lo) * self.strength + hi = 1.0 + (hi - 1.0) * self.strength + gamma = torch.empty((int(sel.sum()), 1, 1, 1), device=device, dtype=dtype).uniform_(lo, hi) + imgs[sel] = imgs[sel].clamp(1e-5, 1.0).pow(gamma) + + # Pequeno desbalanceamento RGB. + sel = choose(self.cfg.get("rgb_gain_p", 0.25)) + if bool(sel.any()): + lo = float(self.cfg.get("rgb_gain_min", 0.94)) + hi = float(self.cfg.get("rgb_gain_max", 1.06)) + lo = 1.0 - (1.0 - lo) * self.strength + hi = 1.0 + (hi - 1.0) * self.strength + gains = torch.empty((int(sel.sum()), C, 1, 1), device=device, dtype=dtype).uniform_(lo, hi) + imgs[sel] = imgs[sel] * gains + + # Sombra direcional suave. + sel = choose(self.cfg.get("shadow_p", 0.35)) + if bool(sel.any()): + yy = torch.linspace(-1.0, 1.0, H, device=device, dtype=dtype).view(H, 1) + xx = torch.linspace(-1.0, 1.0, W, device=device, dtype=dtype).view(1, W) + + for bi in torch.where(sel)[0].tolist(): + angle = random.uniform(0.0, math.pi * 2.0) + proj = math.cos(angle) * xx + math.sin(angle) * yy + proj = (proj - proj.min()) / (proj.max() - proj.min() + 1e-6) + + lo = float(self.cfg.get("shadow_strength_min", 0.12)) + hi = float(self.cfg.get("shadow_strength_max", 0.42)) + strength = random.uniform(lo, hi) * self.strength + + # Metade do gradiente recebe sombra; borda fica suave. + center = random.uniform(0.30, 0.70) + softness = random.uniform(0.12, 0.28) + shade = torch.sigmoid((proj - center) / max(softness, 1e-3)) + factor = 1.0 - strength * shade + imgs[bi] = imgs[bi] * factor.unsqueeze(0) + + # Ruido. + sel = choose(self.cfg.get("noise_p", 0.18)) + if bool(sel.any()): + lo = float(self.cfg.get("noise_sigma_min", 0.003)) + hi = float(self.cfg.get("noise_sigma_max", 0.018)) + sigmas = torch.empty((int(sel.sum()), 1, 1, 1), device=device, dtype=dtype).uniform_(lo, hi) + sigmas = sigmas * self.strength + noise = torch.randn_like(imgs[sel]) * sigmas + imgs[sel] = imgs[sel] + noise + + # Blur leve. + sel = choose(self.cfg.get("blur_p", 0.12)) + if bool(sel.any()): + k = int(self.cfg.get("blur_kernel", 3)) + if k % 2 == 0: + k += 1 + k = max(3, k) + blurred = F.avg_pool2d(imgs[sel], kernel_size=k, stride=1, padding=k // 2) + imgs[sel] = blurred + + # Oclusao opcional, desabilitada por padrao. + sel = choose(self.cfg.get("occlusion_p", 0.0)) + if bool(sel.any()): + fmin = float(self.cfg.get("occlusion_min_frac", 0.04)) + fmax = float(self.cfg.get("occlusion_max_frac", 0.15)) + + for bi in torch.where(sel)[0].tolist(): + frac = random.uniform(fmin, fmax) * self.strength + rh = max(2, int(H * frac)) + rw = max(2, int(W * frac)) + y0 = random.randint(0, max(0, H - rh)) + x0 = random.randint(0, max(0, W - rw)) + factor = random.uniform(0.35, 0.75) + imgs[bi, :, y0:y0 + rh, x0:x0 + rw] *= factor + + return imgs.clamp(0.0, 1.0) + + +# ============================================================================= +# Status head otimizada +# ============================================================================= + +class CorridorStatusHead(nn.Module): + """ + Classificacao global do estado do corredor. + + Diferenca importante para a versao antiga: + - NAO faz upsample da feature profunda ate a resolucao da segmentacao; + - preserva estrutura espacial por pooling 3x4 (configuravel); + - resume separadamente feature profunda e probabilidade semantica; + - concatena vetores compactos antes do MLP. + + Isso e muito mais barato e preserva informacao espacial util para distinguir: + EntrandoRua / CaminhandoRua / SaindoRua / Manobrando / etc. + """ + + def __init__( + self, + feat_ch: int, + num_seg_classes: int, + num_label_classes: int, + pool_hw: Tuple[int, int] = (3, 4), + hidden: int = 256, + dropout: float = 0.20, + detach_seg_summary: bool = True, + seg_summary: str = "probabilities", + ): + super().__init__() + + self.feat_ch = int(feat_ch) + self.num_seg_classes = int(num_seg_classes) + self.num_label_classes = int(num_label_classes) + self.pool_hw = (int(pool_hw[0]), int(pool_hw[1])) + self.detach_seg_summary = bool(detach_seg_summary) + self.seg_summary = str(seg_summary).lower() + + pooled_cells = self.pool_hw[0] * self.pool_hw[1] + in_dim = (self.feat_ch + self.num_seg_classes) * pooled_cells + + self.feat_pool = nn.AdaptiveAvgPool2d(self.pool_hw) + self.seg_pool = nn.AdaptiveAvgPool2d(self.pool_hw) + + self.net = nn.Sequential( + nn.Linear(in_dim, int(hidden)), + nn.ReLU(inplace=True), + nn.Dropout(float(dropout)), + nn.Linear(int(hidden), self.num_label_classes), + ) + + def forward( + self, + feat: torch.Tensor, + seg_logits_native: torch.Tensor, + ) -> torch.Tensor: + feat_grid = self.feat_pool(feat).flatten(1) + + seg_src = seg_logits_native.detach() if self.detach_seg_summary else seg_logits_native + + if self.seg_summary == "probabilities": + seg_src = torch.softmax(seg_src, dim=1) + elif self.seg_summary != "logits": + raise RuntimeError( + f"status_head.seg_summary invalido: {self.seg_summary}. " + "Use 'probabilities' ou 'logits'." + ) + + seg_grid = self.seg_pool(seg_src).flatten(1) + x = torch.cat([feat_grid, seg_grid], dim=1) + return self.net(x) + + def export_config(self) -> dict: + return { + "feat_ch": self.feat_ch, + "num_seg_classes": self.num_seg_classes, + "num_label_classes": self.num_label_classes, + "pool_hw": list(self.pool_hw), + "detach_seg_summary": self.detach_seg_summary, + "seg_summary": self.seg_summary, + } + + +# ============================================================================= +# Pesos de classes e sampler +# ============================================================================= + +def estimate_seg_class_weights( + ds: CorridorDataset, + num_classes: int, + ignore_index: int, + cfg: dict, +) -> torch.Tensor: + """ + Estima frequencia de classes lendo SOMENTE masks. + + A versao v1 fazia ds[idx], carregando imagem inteira para contar pixels + da máscara. Aqui evitamos esse I/O desnecessario. + """ + max_samples = int(cfg.get("seg_max_samples", 1200)) + n = min(len(ds), max_samples) + + if n <= 0: + return torch.ones( + num_classes, + dtype=torch.float32, + ) + + idxs = np.linspace( + 0, + len(ds) - 1, + n, + ).astype(int) + + counts = np.zeros( + num_classes, + dtype=np.float64, + ) + + for idx in idxs: + mask_path = ds._get_mask_path(int(idx)) + + arr = np.asarray( + Image.open(mask_path) + ) + + if arr.ndim == 3: + arr = arr[..., 0] + + mask = arr.reshape(-1).astype(np.int64, copy=False) + mask = mask[mask != ignore_index] + + if mask.size == 0: + continue + + valid = ( + (mask >= 0) + & (mask < num_classes) + ) + + invalid_count = int((~valid).sum()) + + if invalid_count: + bad = np.unique(mask[~valid])[:10].tolist() + raise RuntimeError( + f"Mask normalizada invalida {mask_path}: " + f"{invalid_count} pixels fora das classes; ids={bad}" + ) + + counts += np.bincount( + mask[valid], + minlength=num_classes, + )[:num_classes] + + if counts.sum() <= 0: + return torch.ones( + num_classes, + dtype=torch.float32, + ) + + freq = counts / counts.sum() + freq = np.clip(freq, 1e-12, 1.0) + + method = str( + cfg.get( + "seg_method", + "log_inverse", + ) + ).lower() + + if method == "log_inverse": + offset = float( + cfg.get( + "seg_log_offset", + 1.02, + ) + ) + weights = 1.0 / np.log( + offset + freq + ) + elif method == "sqrt_inverse": + weights = 1.0 / np.sqrt(freq) + else: + raise ValueError( + f"seg_method desconhecido: {method}" + ) + + weights = weights / max( + weights.mean(), + 1e-12, + ) + + weights = np.clip( + weights, + float( + cfg.get( + "seg_min_weight", + 0.35, + ) + ), + float( + cfg.get( + "seg_max_weight", + 3.0, + ) + ), + ) + + print( + f"[LOSS][SEG] counts={counts.astype(np.int64).tolist()}" + ) + print( + f"[LOSS][SEG] freq={np.round(freq, 6).tolist()}" + ) + print( + f"[LOSS][SEG] weights={np.round(weights, 4).tolist()}" + ) + + return torch.tensor( + weights, + dtype=torch.float32, + ) + + +def status_counts(ds: CorridorDataset, num_label_classes: int) -> np.ndarray: + counts = np.zeros(num_label_classes, dtype=np.int64) + for label_id in ds.label_ids: + if 0 <= int(label_id) < num_label_classes: + counts[int(label_id)] += 1 + return counts + + +def estimate_status_class_weights( + ds: CorridorDataset, + num_label_classes: int, + cfg: dict, +) -> torch.Tensor: + counts = status_counts(ds, num_label_classes).astype(np.float64) + + if counts.sum() <= 0: + return torch.ones(num_label_classes, dtype=torch.float32) + + freq = counts / counts.sum() + freq = np.clip(freq, 1e-12, 1.0) + + power = float(cfg.get("status_power", 0.30)) + weights = (1.0 / freq) ** power + weights = weights / max(weights.mean(), 1e-12) + + weights = np.clip( + weights, + float(cfg.get("status_min_weight", 0.50)), + float(cfg.get("status_max_weight", 3.0)), + ) + + print(f"[LOSS][STATUS] counts={counts.astype(np.int64).tolist()}") + print(f"[LOSS][STATUS] weights={np.round(weights, 4).tolist()}") + + return torch.tensor(weights, dtype=torch.float32) + + +def build_training_sampler( + ds: CorridorDataset, + num_label_classes: int, + cfg: dict, +) -> Optional[WeightedRandomSampler]: + """ + Sampler pedagogico. + + mode="joint": + peso considera grupo x estado, estado isolado e grupo isolado. + Potencias baixas impedem que raridades extremas dominem a epoca. + + mode="status": + compatibilidade com a filosofia v1. + + mode="none": + shuffle normal. + """ + mode = str( + cfg.get( + "mode", + "joint", + ) + ).strip().lower() + + if mode in {"none", "off", "false"}: + print("[SAMPLER] disabled") + return None + + labels = [ + int(x) + for x in ds.label_ids + ] + groups = [ + str(x) + for x in ds.group_names + ] + + valid_indices = [ + i + for i, label_id in enumerate(labels) + if 0 <= label_id < num_label_classes + ] + + if not valid_indices: + print("[SAMPLER] nenhum label valido; usando shuffle.") + return None + + status_counts_map = Counter( + labels[i] + for i in valid_indices + ) + + group_counts_map = Counter( + groups[i] + for i in valid_indices + ) + + joint_counts_map = Counter( + (groups[i], labels[i]) + for i in valid_indices + ) + + joint_power = float( + cfg.get( + "joint_power", + 0.30, + ) + ) + status_power = float( + cfg.get( + "status_power", + 0.12, + ) + ) + group_power = float( + cfg.get( + "group_power", + 0.10, + ) + ) + + sample_weights: List[float] = [] + + for i, label_id in enumerate(labels): + if not (0 <= label_id < num_label_classes): + sample_weights.append(1.0) + continue + + group = groups[i] + + if mode == "status": + w = ( + 1.0 + / max( + 1, + status_counts_map[label_id], + ) + ) ** max( + status_power, + joint_power, + ) + elif mode == "joint": + n_joint = max( + 1, + joint_counts_map[ + ( + group, + label_id, + ) + ], + ) + n_status = max( + 1, + status_counts_map[label_id], + ) + n_group = max( + 1, + group_counts_map[group], + ) + + w = ( + (1.0 / n_joint) ** joint_power + * (1.0 / n_status) ** status_power + * (1.0 / n_group) ** group_power + ) + else: + raise ValueError( + f"sampler.mode desconhecido: {mode!r}" + ) + + sample_weights.append( + float(w) + ) + + weights_np = np.asarray( + sample_weights, + dtype=np.float64, + ) + + positive = weights_np[ + weights_np > 0 + ] + + if positive.size == 0: + return None + + weights_np /= positive.mean() + + max_ratio = float( + cfg.get( + "max_weight_ratio", + 5.0, + ) + ) + + positive = weights_np[ + weights_np > 0 + ] + + lo = float( + positive.min() + ) + hi = lo * max( + 1.0, + max_ratio, + ) + + weights_np = np.where( + weights_np > 0, + np.clip( + weights_np, + lo, + hi, + ), + 0.0, + ) + + num_samples = int( + cfg.get( + "samples_per_epoch", + 0, + ) + ) + + if num_samples <= 0: + num_samples = len(ds) + + print(f"[SAMPLER] mode={mode}") + print( + "[SAMPLER] groups=" + + str( + dict( + sorted( + group_counts_map.items() + ) + ) + ) + ) + print( + "[SAMPLER] status=" + + str( + dict( + sorted( + status_counts_map.items() + ) + ) + ) + ) + + if mode == "joint": + print( + "[SAMPLER] joint=" + + str( + { + f"{g}×{lid}": n + for ( + g, + lid, + ), n in sorted( + joint_counts_map.items() + ) + } + ) + ) + + print( + f"[SAMPLER] sample_weight range=" + f"{weights_np.min():.4f}..{weights_np.max():.4f} " + f"samples_per_epoch={num_samples}" + ) + + return WeightedRandomSampler( + weights=torch.tensor( + weights_np, + dtype=torch.double, + ), + num_samples=num_samples, + replacement=True, + ) + + +# ============================================================================= +# Losses +# ============================================================================= + +def boundary_map( + targets: torch.Tensor, + nav_class_id: int, + ignore_index: int, +) -> torch.Tensor: + valid = targets != ignore_index + nav = (targets == nav_class_id).float().unsqueeze(1) + + dil = F.max_pool2d(nav, kernel_size=3, stride=1, padding=1) + ero = -F.max_pool2d(-nav, kernel_size=3, stride=1, padding=1) + + bnd = (dil - ero).squeeze(1) > 1e-6 + return bnd & valid + + +def multiclass_dice_loss( + logits: torch.Tensor, + targets: torch.Tensor, + num_classes: int, + ignore_index: int, + smooth: float = 1.0, +) -> torch.Tensor: + probs = torch.softmax(logits, dim=1) + valid = targets != ignore_index + + losses = [] + + for cid in range(num_classes): + p = probs[:, cid] + t = (targets == cid).float() + + p = p * valid + t = t * valid + + inter = (p * t).sum(dim=(1, 2)) + denom = p.sum(dim=(1, 2)) + t.sum(dim=(1, 2)) + + dice = (2.0 * inter + smooth) / (denom + smooth) + losses.append(1.0 - dice) + + return torch.stack(losses, dim=0).mean() + + +def corridor_seg_loss( + logits: torch.Tensor, + targets: torch.Tensor, + class_weights: Optional[torch.Tensor], + nav_class_id: int, + non_nav_class_id: int, + num_classes: int, + ignore_index: int, + cfg: dict, +) -> Tuple[torch.Tensor, Dict[str, float]]: + ce_map = F.cross_entropy( + logits, + targets, + weight=class_weights, + ignore_index=ignore_index, + reduction="none", + ) + + valid = targets != ignore_index + boundary_boost = float(cfg.get("boundary_boost", 0.15)) + + if boundary_boost > 0: + bnd = boundary_map(targets, nav_class_id, ignore_index) + ce_map = ce_map * (1.0 + boundary_boost * bnd.float()) + + if bool(valid.any()): + ce = ce_map[valid].mean() + else: + ce = logits.sum() * 0.0 + + dice = multiclass_dice_loss( + logits, + targets, + num_classes=num_classes, + ignore_index=ignore_index, + ) + + probs = torch.softmax(logits, dim=1) + p_nav = probs[:, nav_class_id] + gt_non_nav = targets == non_nav_class_id + + if bool(gt_non_nav.any()): + # Quadratico: baixa probabilidade falsa quase nao pesa; alta confianca falsa pesa. + unsafe = (p_nav[gt_non_nav] ** 2).mean() + else: + unsafe = logits.sum() * 0.0 + + w_ce = float(cfg.get("seg_ce_weight", 0.65)) + w_dice = float(cfg.get("seg_dice_weight", 0.35)) + w_unsafe = float(cfg.get("unsafe_nav_weight", 0.06)) + + total = w_ce * ce + w_dice * dice + w_unsafe * unsafe + + return total, { + "ce": float(ce.detach().item()), + "dice": float(dice.detach().item()), + "unsafe": float(unsafe.detach().item()), + } + + +def corridor_status_loss( + logits: torch.Tensor, + labels: torch.Tensor, + class_weights: Optional[torch.Tensor], + label_smoothing: float, +) -> torch.Tensor: + valid = labels != LABEL_IGNORE_INDEX + if not bool(valid.any()): + return logits.sum() * 0.0 + + return F.cross_entropy( + logits[valid], + labels[valid], + weight=class_weights, + label_smoothing=float(label_smoothing), + ) + + +# ============================================================================= +# Metricas +# ============================================================================= + +@torch.no_grad() +def update_confusion_matrix( + cm: torch.Tensor, + preds: torch.Tensor, + labels: torch.Tensor, + num_classes: int, + ignore_index: int, +): + preds = preds.reshape(-1) + labels = labels.reshape(-1) + + valid = labels != ignore_index + preds = preds[valid] + labels = labels[valid] + + if labels.numel() == 0: + return + + idx = labels * num_classes + preds + bins = torch.bincount(idx, minlength=num_classes * num_classes) + cm += bins.view(num_classes, num_classes) + + +@torch.no_grad() +def update_label_confusion( + cm: torch.Tensor, + logits: torch.Tensor, + labels: torch.Tensor, + num_classes: int, +): + valid = labels != LABEL_IGNORE_INDEX + if not bool(valid.any()): + return + + preds = torch.argmax(logits[valid], dim=1) + gt = labels[valid] + + idx = gt * num_classes + preds + bins = torch.bincount(idx, minlength=num_classes * num_classes) + cm += bins.view(num_classes, num_classes) + + +def safe_div(a: torch.Tensor, b: torch.Tensor, eps: float = 1e-8) -> torch.Tensor: + return a / (b + eps) + + +def segmentation_metrics( + cm: torch.Tensor, + nav_class_id: int, + non_nav_class_id: int, +) -> Dict[str, Any]: + cmf = cm.float() + tp = torch.diag(cmf) + fp = cmf.sum(0) - tp + fn = cmf.sum(1) - tp + + iou = safe_div(tp, tp + fp + fn) + precision = safe_div(tp, tp + fp) + recall = safe_div(tp, tp + fn) + f1 = safe_div(2 * precision * recall, precision + recall) + + total = cmf.sum() + acc = safe_div(tp.sum(), total) + + # Safety: entre pixels realmente nao navegaveis, quantos foram liberados como navegaveis? + unsafe_nav = cmf[non_nav_class_id, nav_class_id] + gt_non_nav = cmf[non_nav_class_id].sum() + unsafe_nav_rate = safe_div(unsafe_nav, gt_non_nav) + safety = 1.0 - unsafe_nav_rate + + return { + "acc": float(acc.item()), + "miou": float(iou.mean().item()), + "iou_per_class": iou.cpu().tolist(), + "precision_per_class": precision.cpu().tolist(), + "recall_per_class": recall.cpu().tolist(), + "f1_per_class": f1.cpu().tolist(), + "nav_iou": float(iou[nav_class_id].item()), + "nav_precision": float(precision[nav_class_id].item()), + "nav_recall": float(recall[nav_class_id].item()), + "nav_f1": float(f1[nav_class_id].item()), + "non_nav_iou": float(iou[non_nav_class_id].item()), + "unsafe_nav_rate": float(unsafe_nav_rate.item()), + "safety": float(safety.item()), + "support_per_class": cmf.sum(1).cpu().tolist(), + "gt_nav_pixels": int(cmf[nav_class_id].sum().item()), + "gt_non_nav_pixels": int(gt_non_nav.item()), + } + + +def status_metrics(cm: torch.Tensor) -> Dict[str, Any]: + cmf = cm.float() + tp = torch.diag(cmf) + fp = cmf.sum(0) - tp + fn = cmf.sum(1) - tp + support = cmf.sum(1) + + precision = safe_div(tp, tp + fp) + recall = safe_div(tp, tp + fn) + f1 = safe_div(2 * precision * recall, precision + recall) + + acc = safe_div(tp.sum(), cmf.sum()) + + present = support > 0 + if bool(present.any()): + macro_f1 = f1[present].mean() + macro_recall = recall[present].mean() + else: + macro_f1 = torch.tensor(0.0, device=cm.device) + macro_recall = torch.tensor(0.0, device=cm.device) + + return { + "acc": float(acc.item()), + "macro_f1": float(macro_f1.item()), + "macro_recall": float(macro_recall.item()), + "precision_per_class": precision.cpu().tolist(), + "recall_per_class": recall.cpu().tolist(), + "f1_per_class": f1.cpu().tolist(), + "support_per_class": support.cpu().tolist(), + } + + +def field_score(seg: dict, status: dict, cfg: dict) -> float: + weights = { + "nav_iou": float(cfg.get("nav_iou", 0.40)), + "nav_f1": float(cfg.get("nav_f1", 0.20)), + "status_macro_f1": float(cfg.get("status_macro_f1", 0.20)), + "safety": float(cfg.get("safety", 0.20)), + } + + total_w = sum(weights.values()) + if total_w <= 0: + raise ValueError("score weights precisam somar > 0") + + score = ( + weights["nav_iou"] * seg["nav_iou"] + + weights["nav_f1"] * seg["nav_f1"] + + weights["status_macro_f1"] * status["macro_f1"] + + weights["safety"] * seg["safety"] + ) / total_w + + return float(score) + + + +def evaluate_operational_gate( + *, + epoch: int, + val_metrics: dict, + cfg: dict, + label_name_by_id: Dict[int, str], +) -> Dict[str, Any]: + """Avalia se uma epoca e elegivel para concorrer a best_operational.""" + enabled = bool(cfg.get("enabled", True)) + + result: Dict[str, Any] = { + "schema": "agrobot.corridor.operational_gate.v1", + "enabled": enabled, + "eligible": True, + "checks": {}, + "failures": [], + "warnings": [], + } + + if not enabled: + result["warnings"].append("GATE_DISABLED") + return result + + seg = val_metrics["seg"] + status = val_metrics["status"] + seg_by_group = val_metrics.get("seg_by_group", {}) or {} + + def add_min(name: str, value: float, threshold_value) -> None: + if threshold_value is None: + return + threshold = float(threshold_value) + value_f = float(value) + passed = value_f >= threshold + result["checks"][name] = { + "passed": bool(passed), + "value": value_f, + "operator": ">=", + "threshold": threshold, + } + if not passed: + result["failures"].append( + f"{name}:{value_f:.6f}<{threshold:.6f}" + ) + + def add_max(name: str, value: float, threshold_value) -> None: + if threshold_value is None: + return + threshold = float(threshold_value) + value_f = float(value) + passed = value_f <= threshold + result["checks"][name] = { + "passed": bool(passed), + "value": value_f, + "operator": "<=", + "threshold": threshold, + } + if not passed: + result["failures"].append( + f"{name}:{value_f:.6f}>{threshold:.6f}" + ) + + min_epoch = int(cfg.get("min_epoch", 1)) + epoch_passed = int(epoch) >= min_epoch + result["checks"]["min_epoch"] = { + "passed": bool(epoch_passed), + "value": int(epoch), + "operator": ">=", + "threshold": int(min_epoch), + } + if not epoch_passed: + result["failures"].append( + f"min_epoch:{int(epoch)}<{int(min_epoch)}" + ) + + add_min("min_nav_iou", seg["nav_iou"], cfg.get("min_nav_iou")) + add_min("min_nav_f1", seg["nav_f1"], cfg.get("min_nav_f1")) + add_min("min_non_nav_iou", seg["non_nav_iou"], cfg.get("min_non_nav_iou")) + add_min( + "min_status_macro_f1", + status["macro_f1"], + cfg.get("min_status_macro_f1"), + ) + add_max( + "max_unsafe_nav_rate", + seg["unsafe_nav_rate"], + cfg.get("max_unsafe_nav_rate"), + ) + + group_cfg = cfg.get("group_unsafe", {}) or {} + if bool(group_cfg.get("enabled", True)): + min_non_nav_pixels = int(group_cfg.get("min_non_nav_pixels", 0)) + default_max_rate = float(group_cfg.get("max_rate", 1.0)) + by_group = group_cfg.get("max_rate_by_group", {}) or {} + group_checks = {} + + for group_name in sorted(seg_by_group): + gm = seg_by_group[group_name] + support = int(gm.get("gt_non_nav_pixels", 0)) + rate = float(gm.get("unsafe_nav_rate", 0.0)) + threshold = float(by_group.get(group_name, default_max_rate)) + applicable = support >= min_non_nav_pixels + passed = (rate <= threshold) if applicable else True + + group_checks[group_name] = { + "applicable": bool(applicable), + "passed": bool(passed), + "gt_non_nav_pixels": int(support), + "min_non_nav_pixels": int(min_non_nav_pixels), + "unsafe_nav_rate": rate, + "max_unsafe_nav_rate": threshold, + } + + if applicable and not passed: + result["failures"].append( + f"group_unsafe:{group_name}:{rate:.6f}>{threshold:.6f}" + ) + + result["checks"]["group_unsafe"] = group_checks + + status_cfg = cfg.get("status_recall", {}) or {} + if bool(status_cfg.get("enabled", False)): + minimums = status_cfg.get("minimums", {}) or {} + min_support = int(status_cfg.get("min_support", 1)) + recalls = status.get("recall_per_class", []) + supports = status.get("support_per_class", []) + status_checks = {} + name_to_id = { + str(name): int(cid) + for cid, name in label_name_by_id.items() + } + + for class_name, threshold_raw in minimums.items(): + if class_name not in name_to_id: + status_checks[class_name] = { + "applicable": False, + "passed": False, + "reason": "UNKNOWN_STATUS_CLASS", + } + result["failures"].append( + f"status_recall:{class_name}:UNKNOWN_STATUS_CLASS" + ) + continue + + cid = name_to_id[class_name] + support = int(supports[cid] if cid < len(supports) else 0) + recall = float(recalls[cid] if cid < len(recalls) else 0.0) + threshold = float(threshold_raw) + applicable = support >= min_support + passed = (recall >= threshold) if applicable else True + + status_checks[class_name] = { + "applicable": bool(applicable), + "passed": bool(passed), + "support": int(support), + "min_support": int(min_support), + "recall": recall, + "min_recall": threshold, + } + + if applicable and not passed: + result["failures"].append( + f"status_recall:{class_name}:{recall:.6f}<{threshold:.6f}" + ) + if not applicable: + result["warnings"].append( + f"status_recall:{class_name}:support={support} str: + if not bool(gate.get("enabled", True)): + return "DISABLED" + if bool(gate.get("eligible", False)): + return "ELIGIBLE" + + failures = list(gate.get("failures", [])) + shown = failures[:max_failures] + suffix = ( + "" + if len(failures) <= max_failures + else f" (+{len(failures) - max_failures})" + ) + return "FAIL | " + " | ".join(shown) + suffix + + +def pretty_status(values: Sequence[float], names: Dict[int, str]) -> str: + return " | ".join( + f"{names.get(i, str(i))}:{float(v):.3f}" + for i, v in enumerate(values) + ) + + +# ============================================================================= +# Optimizer / scheduler +# ============================================================================= + +def is_no_decay_param(name: str, p: torch.Tensor, cfg: dict) -> bool: + lname = name.lower() + + if bool(cfg.get("no_decay_bias", True)) and lname.endswith(".bias"): + return True + + if bool(cfg.get("no_decay_norm", True)): + if p.ndim <= 1: + return True + if any(token in lname for token in ("norm", "bn", "batchnorm", "layernorm")): + return True + + return False + + +def append_param_groups( + groups: List[dict], + named_params: Sequence[Tuple[str, torch.nn.Parameter]], + lr: float, + wd: float, + prefix: str, + cfg: dict, +): + decay, no_decay = [], [] + + for name, p in named_params: + if not p.requires_grad: + # Mesmo congelado agora, pode ser descongelado futuramente. + # Mantemos no optimizer. + pass + + full_name = f"{prefix}.{name}" + + if is_no_decay_param(full_name, p, cfg): + no_decay.append(p) + else: + decay.append(p) + + if decay: + groups.append({ + "params": decay, + "lr": float(lr), + "weight_decay": float(wd), + "group_name": f"{prefix}_decay", + }) + + if no_decay: + groups.append({ + "params": no_decay, + "lr": float(lr), + "weight_decay": 0.0, + "group_name": f"{prefix}_no_decay", + }) + + +def build_optimizer( + base_model: SegformerForSemanticSegmentation, + status_head: CorridorStatusHead, + base_lr: float, + default_wd: float, + cfg: dict, +) -> torch.optim.Optimizer: + wd = cfg.get("weight_decay") + wd = float(default_wd if wd is None else wd) + + enc_lr = float(base_lr) * float(cfg.get("encoder_lr_mult", 0.50)) + dec_lr = float(base_lr) * float(cfg.get("decoder_lr_mult", 1.00)) + status_lr = float(base_lr) * float(cfg.get("status_lr_mult", 2.00)) + + groups: List[dict] = [] + + append_param_groups( + groups, + list(base_model.segformer.named_parameters()), + enc_lr, + wd, + "encoder", + cfg, + ) + + append_param_groups( + groups, + list(base_model.decode_head.named_parameters()), + dec_lr, + wd, + "decoder", + cfg, + ) + + append_param_groups( + groups, + list(status_head.named_parameters()), + status_lr, + wd, + "status", + cfg, + ) + + betas = cfg.get("betas", [0.9, 0.999]) + eps = float(cfg.get("eps", 1e-8)) + + opt = torch.optim.AdamW( + groups, + betas=(float(betas[0]), float(betas[1])), + eps=eps, + ) + + print( + f"[OPT] lr encoder={enc_lr:.3e} decoder={dec_lr:.3e} " + f"status={status_lr:.3e} wd={wd:.3e}" + ) + + return opt + + +def build_scheduler( + optimizer: torch.optim.Optimizer, + total_steps: int, + cfg: dict, +): + mode = str(cfg.get("mode", "poly")).lower() + + warmup_ratio = float(cfg.get("warmup_ratio", 0.05)) + warmup_steps = int(round(total_steps * warmup_ratio)) + warmup_start = float(cfg.get("warmup_start_factor", 0.10)) + power = float(cfg.get("poly_power", 1.0)) + min_lr_ratio = float(cfg.get("min_lr_ratio", 0.02)) + + def lr_lambda(step: int): + if warmup_steps > 0 and step < warmup_steps: + alpha = step / max(1, warmup_steps) + return warmup_start + alpha * (1.0 - warmup_start) + + progress = (step - warmup_steps) / max(1, total_steps - warmup_steps) + progress = min(max(progress, 0.0), 1.0) + + if mode == "poly": + factor = (1.0 - progress) ** power + elif mode == "cosine": + factor = 0.5 * (1.0 + math.cos(math.pi * progress)) + elif mode == "constant": + factor = 1.0 + else: + raise ValueError(f"scheduler mode desconhecido: {mode}") + + return min_lr_ratio + (1.0 - min_lr_ratio) * factor + + return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lr_lambda) + + +def set_encoder_trainable( + base_model: SegformerForSemanticSegmentation, + trainable: bool, +): + for p in base_model.segformer.parameters(): + p.requires_grad = bool(trainable) + + +# ============================================================================= +# Training history +# ============================================================================= + +def _history_json_safe(value): + if isinstance(value, dict): + return {str(k): _history_json_safe(v) for k, v in value.items()} + if isinstance(value, (list, tuple)): + return [_history_json_safe(v) for v in value] + if isinstance(value, Path): + return str(value) + if isinstance(value, np.ndarray): + return value.tolist() + if isinstance(value, np.generic): + return value.item() + if torch.is_tensor(value): + if value.numel() == 1: + return value.detach().cpu().item() + return value.detach().cpu().tolist() + if isinstance(value, (str, int, float, bool)) or value is None: + return value + return str(value) + + +def _flatten_history(value, prefix: str = "", out: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: + if out is None: + out = {} + if isinstance(value, dict): + for key, item in value.items(): + child = f"{prefix}.{key}" if prefix else str(key) + _flatten_history(item, child, out) + return out + if isinstance(value, (list, tuple)): + for i, item in enumerate(value): + child = f"{prefix}.{i}" if prefix else str(i) + _flatten_history(item, child, out) + if not value and prefix: + out[prefix] = "" + return out + out[prefix] = value + return out + + +def _load_history_jsonl(jsonl_path: Path) -> List[dict]: + if not jsonl_path.is_file(): + return [] + rows = [] + with jsonl_path.open("r", encoding="utf-8") as f: + for lineno, raw in enumerate(f, 1): + line = raw.strip() + if not line: + continue + try: + item = json.loads(line) + except Exception as exc: + raise RuntimeError( + f"Historico JSONL corrompido em {jsonl_path}:{lineno}: {exc}" + ) from exc + if not isinstance(item, dict): + raise RuntimeError(f"Historico JSONL invalido em {jsonl_path}:{lineno}") + rows.append(item) + return rows + + +def _rewrite_history_files(*, records: Sequence[dict], jsonl_path: Path, json_path: Path, csv_path: Path) -> None: + records_safe = [_history_json_safe(r) for r in records] + records_safe = sorted(records_safe, key=lambda r: int(r.get("epoch", 0))) + jsonl_path.parent.mkdir(parents=True, exist_ok=True) + + tmp_jsonl = jsonl_path.with_suffix(jsonl_path.suffix + ".tmp") + with tmp_jsonl.open("w", encoding="utf-8") as f: + for record in records_safe: + f.write(json.dumps(record, ensure_ascii=False, separators=(",", ":"))) + f.write("\n") + os.replace(tmp_jsonl, jsonl_path) + + tmp_json = json_path.with_suffix(json_path.suffix + ".tmp") + with tmp_json.open("w", encoding="utf-8") as f: + json.dump(records_safe, f, ensure_ascii=False, indent=2) + os.replace(tmp_json, json_path) + + flat_rows = [_flatten_history(record) for record in records_safe] + fieldnames = [] + seen = set() + for row in flat_rows: + for key in row: + if key not in seen: + seen.add(key) + fieldnames.append(key) + + priority = [ + "epoch", "created_at", "field_score", + "operational_eligible", "operational_score_improved", + "encoder_trainable", + "augmentation_strength", "status_weight", + "train.loss", "train.loss_seg", "train.loss_status", + "train.seg.nav_iou", "train.seg.nav_f1", "train.status.macro_f1", + "val.loss", "val.loss_seg", "val.loss_status", + "val.seg.nav_iou", "val.seg.non_nav_iou", + "val.seg.nav_precision", "val.seg.nav_recall", "val.seg.nav_f1", + "val.seg.unsafe_nav_rate", "val.seg.safety", + "val.status.acc", "val.status.macro_f1", "val.status.macro_recall", + "epoch_time_s", + ] + ordered = [key for key in priority if key in seen] + ordered += [key for key in fieldnames if key not in ordered] + + tmp_csv = csv_path.with_suffix(csv_path.suffix + ".tmp") + with tmp_csv.open("w", newline="", encoding="utf-8-sig") as f: + writer = csv.DictWriter(f, fieldnames=ordered, extrasaction="ignore") + writer.writeheader() + for row in flat_rows: + writer.writerow(row) + os.replace(tmp_csv, csv_path) + + +def initialize_training_history(*, save_dir: Path, resume: bool, start_epoch: int) -> Tuple[Path, Path, Path]: + jsonl_path = save_dir / "metrics_history.jsonl" + json_path = save_dir / "metrics_history.json" + csv_path = save_dir / "metrics_history.csv" + + if not resume: + for p in (jsonl_path, json_path, csv_path): + if p.exists(): + p.unlink() + return jsonl_path, json_path, csv_path + + old = _load_history_jsonl(jsonl_path) + kept = [row for row in old if int(row.get("epoch", 0)) < int(start_epoch)] + if old: + _rewrite_history_files( + records=kept, + jsonl_path=jsonl_path, + json_path=json_path, + csv_path=csv_path, + ) + return jsonl_path, json_path, csv_path + + +def save_epoch_history(*, record: dict, jsonl_path: Path, json_path: Path, csv_path: Path) -> None: + records = _load_history_jsonl(jsonl_path) + epoch = int(record["epoch"]) + records = [row for row in records if int(row.get("epoch", -1)) != epoch] + records.append(_history_json_safe(record)) + _rewrite_history_files( + records=records, + jsonl_path=jsonl_path, + json_path=json_path, + csv_path=csv_path, + ) + + +# ============================================================================= +# Checkpoint +# ============================================================================= + +@dataclass +class EarlyStopState: + best: float = -1e9 + bad_epochs: int = 0 + + +def save_checkpoint( + path: Path, + base_model: SegformerForSemanticSegmentation, + status_head: CorridorStatusHead, + optimizer: torch.optim.Optimizer, + scheduler, + scaler: GradScaler, + epoch: int, + bests: dict, + early_stop: EarlyStopState, + metadata: dict, +): + payload = { + "trainer_version": TRAINER_VERSION, + "epoch": int(epoch), + "model": base_model.state_dict(), + "status_head": status_head.state_dict(), + "optimizer": optimizer.state_dict(), + "scheduler": scheduler.state_dict() if scheduler is not None else None, + "scaler": scaler.state_dict() if scaler is not None else None, + "bests": dict(bests), + "early_stop": { + "best": float(early_stop.best), + "bad_epochs": int(early_stop.bad_epochs), + }, + "rng_state": capture_rng_state(), + "metadata": metadata, + } + + torch.save(payload, path) + + +def load_checkpoint( + path: Path, + base_model: SegformerForSemanticSegmentation, + status_head: CorridorStatusHead, + optimizer: Optional[torch.optim.Optimizer] = None, + scheduler=None, + scaler: Optional[GradScaler] = None, +) -> dict: + ckpt = torch.load(path, map_location="cpu", weights_only=False) + + base_model.load_state_dict(ckpt["model"], strict=True) + status_head.load_state_dict(ckpt["status_head"], strict=True) + + if optimizer is not None and ckpt.get("optimizer") is not None: + optimizer.load_state_dict(ckpt["optimizer"]) + + if scheduler is not None and ckpt.get("scheduler") is not None: + scheduler.load_state_dict(ckpt["scheduler"]) + + if scaler is not None and ckpt.get("scaler") is not None: + scaler.load_state_dict(ckpt["scaler"]) + + restore_rng_state(ckpt.get("rng_state")) + return ckpt + + +# ============================================================================= +# Epoch +# ============================================================================= + +def current_status_weight(epoch: int, loss_cfg: dict) -> float: + target = float(loss_cfg.get("status_weight", 0.40)) + ramp = int(loss_cfg.get("status_ramp_epochs", 8)) + + if ramp <= 0: + return target + + alpha = min(1.0, max(0.10, float(epoch) / float(ramp))) + return target * alpha + + +def run_epoch( + *, + base_model: SegformerForSemanticSegmentation, + status_head: CorridorStatusHead, + loader: DataLoader, + device: torch.device, + mean: torch.Tensor, + std: torch.Tensor, + num_seg_classes: int, + num_label_classes: int, + num_groups: int, + group_name_by_id: Dict[int, str], + nav_class_id: int, + non_nav_class_id: int, + seg_class_weights: Optional[torch.Tensor], + status_class_weights: Optional[torch.Tensor], + loss_cfg: dict, + augmenter: Optional[FieldCorridorAugmenter], + optimizer: Optional[torch.optim.Optimizer], + scheduler, + scaler: Optional[GradScaler], + amp: bool, + train: bool, + grad_accum: int, + grad_clip_norm: float, + epoch: int, +) -> Dict[str, Any]: + base_model.train(train) + status_head.train(train) + + if augmenter is not None: + augmenter.set_epoch(epoch) + + cm_seg = torch.zeros( + (num_seg_classes, num_seg_classes), + dtype=torch.int64, + device=device, + ) + cm_status = torch.zeros( + (num_label_classes, num_label_classes), + dtype=torch.int64, + device=device, + ) + + cm_seg_group = torch.zeros( + ( + num_groups, + num_seg_classes, + num_seg_classes, + ), + dtype=torch.int64, + device=device, + ) + + cm_status_group = torch.zeros( + ( + num_groups, + num_label_classes, + num_label_classes, + ), + dtype=torch.int64, + device=device, + ) + + total_loss = 0.0 + total_seg = 0.0 + total_status = 0.0 + total_ce = 0.0 + total_dice = 0.0 + total_unsafe = 0.0 + n_batches = 0 + + status_w = current_status_weight(epoch, loss_cfg) + t0 = time.perf_counter() + + if train and optimizer is not None: + optimizer.zero_grad(set_to_none=True) + + for step, (imgs, masks, labels, group_ids) in enumerate(loader, start=1): + imgs = imgs.to(device, non_blocking=True) + masks = masks.to(device, non_blocking=True) + labels = labels.to(device, non_blocking=True) + group_ids = group_ids.to(device, non_blocking=True) + + if train and augmenter is not None: + imgs, masks = augmenter(imgs, masks) + + imgs_norm = (imgs - mean) / std + + with torch.set_grad_enabled(train): + with torch.autocast( + device_type=device.type, + enabled=bool(amp and device.type == "cuda"), + ): + out = base_model( + pixel_values=imgs_norm, + output_hidden_states=True, + return_dict=True, + ) + + seg_logits_native = out.logits + feat = out.hidden_states[-1] + + status_logits = status_head(feat, seg_logits_native) + + seg_logits = seg_logits_native + if seg_logits.shape[-2:] != masks.shape[-2:]: + seg_logits = F.interpolate( + seg_logits, + size=masks.shape[-2:], + mode="bilinear", + align_corners=False, + ) + + seg_loss, seg_parts = corridor_seg_loss( + seg_logits, + masks, + class_weights=seg_class_weights, + nav_class_id=nav_class_id, + non_nav_class_id=non_nav_class_id, + num_classes=num_seg_classes, + ignore_index=IGNORE_INDEX, + cfg=loss_cfg, + ) + + status_loss = corridor_status_loss( + status_logits, + labels, + class_weights=status_class_weights, + label_smoothing=float(loss_cfg.get("status_label_smoothing", 0.03)), + ) + + loss = seg_loss + status_w * status_loss + + loss_for_backward = loss / max(1, grad_accum) + + if train and optimizer is not None: + assert scaler is not None + + scaler.scale(loss_for_backward).backward() + + should_step = ( + step % max(1, grad_accum) == 0 + or step == len(loader) + ) + + if should_step: + scaler.unscale_(optimizer) + + if grad_clip_norm > 0: + params = list(base_model.parameters()) + list(status_head.parameters()) + torch.nn.utils.clip_grad_norm_(params, max_norm=grad_clip_norm) + + scaler.step(optimizer) + scaler.update() + optimizer.zero_grad(set_to_none=True) + + if scheduler is not None: + scheduler.step() + + total_loss += float(loss.detach().item()) + total_seg += float(seg_loss.detach().item()) + total_status += float(status_loss.detach().item()) + total_ce += seg_parts["ce"] + total_dice += seg_parts["dice"] + total_unsafe += seg_parts["unsafe"] + n_batches += 1 + + with torch.no_grad(): + pred_seg = torch.argmax(seg_logits, dim=1) + update_confusion_matrix( + cm_seg, + pred_seg, + masks, + num_seg_classes, + IGNORE_INDEX, + ) + update_label_confusion( + cm_status, + status_logits, + labels, + num_label_classes, + ) + + for gid in range(num_groups): + sel = group_ids == gid + + if not bool(sel.any()): + continue + + update_confusion_matrix( + cm_seg_group[gid], + pred_seg[sel], + masks[sel], + num_seg_classes, + IGNORE_INDEX, + ) + + update_label_confusion( + cm_status_group[gid], + status_logits[sel], + labels[sel], + num_label_classes, + ) + + seg_metrics = segmentation_metrics( + cm_seg, + nav_class_id=nav_class_id, + non_nav_class_id=non_nav_class_id, + ) + stat_metrics = status_metrics(cm_status) + + seg_by_group = {} + status_by_group = {} + + for gid in range(num_groups): + gname = group_name_by_id.get( + gid, + str(gid), + ) + + seg_by_group[gname] = segmentation_metrics( + cm_seg_group[gid], + nav_class_id=nav_class_id, + non_nav_class_id=non_nav_class_id, + ) + + status_by_group[gname] = status_metrics( + cm_status_group[gid] + ) + + return { + "loss": total_loss / max(1, n_batches), + "loss_seg": total_seg / max(1, n_batches), + "loss_status": total_status / max(1, n_batches), + "loss_ce": total_ce / max(1, n_batches), + "loss_dice": total_dice / max(1, n_batches), + "loss_unsafe": total_unsafe / max(1, n_batches), + "status_weight": status_w, + "seg": seg_metrics, + "status": stat_metrics, + "seg_by_group": seg_by_group, + "status_by_group": status_by_group, + "time_s": time.perf_counter() - t0, + } + + + +# ============================================================================= +# Pipeline contract / preflight +# ============================================================================= + +def _stem_set(folder: Path, exts: Tuple[str, ...]) -> set[str]: + if not folder.is_dir(): + return set() + + return { + p.stem + for p in folder.iterdir() + if p.is_file() + and p.suffix.lower() in exts + } + + +def audit_split_root( + split_root: Path, + expected_wh: Tuple[int, int], + num_seg_classes: int, + nav_class_id: int, + label_classes: Sequence[str], + max_samples: int = 0, +) -> dict: + """ + Auditoria forte antes de entregar qualquer batch ao modelo. + """ + images_dir = split_root / "images" + masks_dir = split_root / "masks" + labels_dir = split_root / "labels" + + for folder in ( + images_dir, + masks_dir, + labels_dir, + ): + if not folder.is_dir(): + raise FileNotFoundError( + f"Split incompleto, pasta ausente: {folder}" + ) + + image_stems = _stem_set( + images_dir, + ( + ".png", + ".jpg", + ".jpeg", + ".bmp", + ".webp", + ), + ) + mask_stems = _stem_set( + masks_dir, + (".png",), + ) + label_stems = _stem_set( + labels_dir, + (".json",), + ) + + if not image_stems: + raise RuntimeError( + f"Split vazio: {split_root}" + ) + + if not ( + image_stems + == mask_stems + == label_stems + ): + missing_mask = sorted( + image_stems - mask_stems + )[:10] + missing_label = sorted( + image_stems - label_stems + )[:10] + orphan_mask = sorted( + mask_stems - image_stems + )[:10] + orphan_label = sorted( + label_stems - image_stems + )[:10] + + raise RuntimeError( + f"Contrato quebrado em {split_root}.\n" + f"missing_mask={missing_mask}\n" + f"missing_label={missing_label}\n" + f"orphan_mask={orphan_mask}\n" + f"orphan_label={orphan_label}" + ) + + ordered = sorted( + image_stems + ) + + if max_samples > 0 and len(ordered) > max_samples: + idxs = np.linspace( + 0, + len(ordered) - 1, + max_samples, + ).astype(int) + audit_stems = [ + ordered[int(i)] + for i in idxs + ] + else: + audit_stems = ordered + + W, H = int(expected_wh[0]), int(expected_wh[1]) + + status_counts_local = Counter() + group_counts_local = Counter() + joint_counts_local = Counter() + nav_pct_values = [] + + for stem in audit_stems: + # Imagem. + image_candidates = [ + p + for p in images_dir.glob(stem + ".*") + if p.suffix.lower() + in { + ".png", + ".jpg", + ".jpeg", + ".bmp", + ".webp", + } + ] + + if len(image_candidates) != 1: + raise RuntimeError( + f"Esperava 1 imagem para {stem}, encontrei {image_candidates}" + ) + + img = cv2.imread( + str(image_candidates[0]), + cv2.IMREAD_COLOR, + ) + + if img is None: + raise RuntimeError( + f"Nao consegui ler imagem: {image_candidates[0]}" + ) + + if img.shape[:2] != ( + H, + W, + ): + raise RuntimeError( + f"Imagem {stem} shape={img.shape[:2]}, esperado={(H,W)}" + ) + + # Mask ID. + mask_path = masks_dir / f"{stem}.png" + mask = cv2.imread( + str(mask_path), + cv2.IMREAD_UNCHANGED, + ) + + if mask is None: + raise RuntimeError( + f"Nao consegui ler mask: {mask_path}" + ) + + if mask.ndim != 2: + raise RuntimeError( + f"Mask normalizada precisa ser 1 canal IDs: " + f"{mask_path} shape={mask.shape}" + ) + + if mask.shape != ( + H, + W, + ): + raise RuntimeError( + f"Mask {stem} shape={mask.shape}, esperado={(H,W)}" + ) + + ids = sorted( + int(x) + for x in np.unique(mask) + ) + + if not set(ids).issubset( + set( + range( + num_seg_classes + ) + ) + ): + raise RuntimeError( + f"Mask {stem} possui IDs invalidos={ids}" + ) + + # Label. + label_path = labels_dir / f"{stem}.json" + + with label_path.open( + "r", + encoding="utf-8", + ) as f: + data = json.load(f) + + label_id = int( + data.get( + "label_id", + -999, + ) + ) + + if not ( + 0 + <= label_id + < len(label_classes) + ): + raise RuntimeError( + f"label_id invalido em {label_path}: {label_id}" + ) + + expected_name = str( + label_classes[label_id] + ) + got_name = str( + data.get( + "estado_corredor", + expected_name, + ) + ) + + if got_name != expected_name: + raise RuntimeError( + f"Label inconsistente em {label_path}: " + f"id={label_id} => {expected_name}, arquivo={got_name}" + ) + + group = str( + data.get( + "split_source_group", + data.get( + "group", + "unknown", + ), + ) + ) + + status_counts_local[label_id] += 1 + group_counts_local[group] += 1 + joint_counts_local[ + ( + group, + label_id, + ) + ] += 1 + + # nav_class_id nao e necessario aqui; para auditoria de densidade + # basta guardar composicao binaria pelo maior ID se houver 2 classes. + if num_seg_classes == 2: + nav_pct_values.append( + float( + ( + mask == int(nav_class_id) + ).mean() + * 100.0 + ) + ) + + return { + "count_total": len(image_stems), + "count_audited": len(audit_stems), + "status_counts_audited": { + str(k): int(v) + for k, v in sorted( + status_counts_local.items() + ) + }, + "group_counts_audited": { + str(k): int(v) + for k, v in sorted( + group_counts_local.items() + ) + }, + "joint_counts_audited": { + f"{g}__label_{lid}": int(v) + for ( + g, + lid, + ), v in sorted( + joint_counts_local.items() + ) + }, + "nav_pct_audited_min": ( + float(min(nav_pct_values)) + if nav_pct_values + else None + ), + "nav_pct_audited_max": ( + float(max(nav_pct_values)) + if nav_pct_values + else None + ), + } + + +def audit_pipeline_contract( + resolution_root: Path, + train_root: Path, + val_root: Path, + expected_wh: Tuple[int, int], + num_seg_classes: int, + nav_class_id: int, + label_classes: Sequence[str], + strict_contract: bool, + max_samples: int, +) -> dict: + split_report_path = ( + resolution_root + / "split_report.json" + ) + manifest_path = ( + resolution_root + / "split_manifest.csv" + ) + + if strict_contract: + for required in ( + split_report_path, + manifest_path, + resolution_root / "norm_stats.json", + ): + if not required.is_file(): + raise FileNotFoundError( + f"Artefato obrigatorio da etapa split ausente: {required}" + ) + + train_audit = audit_split_root( + train_root, + expected_wh=expected_wh, + num_seg_classes=num_seg_classes, + nav_class_id=nav_class_id, + label_classes=label_classes, + max_samples=max_samples, + ) + + val_audit = audit_split_root( + val_root, + expected_wh=expected_wh, + num_seg_classes=num_seg_classes, + nav_class_id=nav_class_id, + label_classes=label_classes, + max_samples=max_samples, + ) + + train_stems = _stem_set( + train_root / "images", + ( + ".png", + ".jpg", + ".jpeg", + ".bmp", + ".webp", + ), + ) + val_stems = _stem_set( + val_root / "images", + ( + ".png", + ".jpg", + ".jpeg", + ".bmp", + ".webp", + ), + ) + + overlap = sorted( + train_stems + & val_stems + ) + + if overlap: + raise RuntimeError( + f"Vazamento por stem entre train e val: {overlap[:20]}" + ) + + split_report = {} + + if split_report_path.is_file(): + with split_report_path.open( + "r", + encoding="utf-8", + ) as f: + split_report = json.load(f) + + report_wh = split_report.get( + "resolution_wh" + ) + + if ( + strict_contract + and report_wh is not None + and [ + int(report_wh[0]), + int(report_wh[1]), + ] + != [ + int(expected_wh[0]), + int(expected_wh[1]), + ] + ): + raise RuntimeError( + f"split_report resolution={report_wh} " + f"mas config={list(expected_wh)}" + ) + + if strict_contract: + tr_report = int( + split_report.get( + "train_count", + -1, + ) + ) + va_report = int( + split_report.get( + "val_count", + -1, + ) + ) + + if ( + tr_report + != train_audit[ + "count_total" + ] + or va_report + != val_audit[ + "count_total" + ] + ): + raise RuntimeError( + "Contagem do split_report diverge das pastas atuais: " + f"report train/val={tr_report}/{va_report}, " + f"pastas={train_audit['count_total']}/{val_audit['count_total']}" + ) + + return { + "train": train_audit, + "val": val_audit, + "split_report_path": str( + split_report_path + ), + "split_manifest_path": str( + manifest_path + ), + "split_report": split_report, + } + + +# ============================================================================= +# Main +# ============================================================================= + +def main(): + ap = argparse.ArgumentParser( + description="Agri Corridor Teacher v2 - SegFormer OAK-D Lite" + ) + + ap.add_argument("--config", default="config.json") + + ap.add_argument("--epochs", type=int, default=140) + ap.add_argument("--batch", type=int, default=4) + ap.add_argument("--num_workers", type=int, default=4) + ap.add_argument("--seed", type=int, default=42) + + ap.add_argument("--lr", type=float, default=3e-5) + ap.add_argument("--wd", type=float, default=1e-2) + ap.add_argument("--grad_accum", type=int, default=2) + + ap.add_argument("--amp", action="store_true") + ap.add_argument("--amp_val", action="store_true") + ap.add_argument( + "--grad_ckpt", + action="store_true", + help="Ativa gradient checkpointing no SegFormer se suportado.", + ) + + ap.add_argument( + "--norm_stats", + type=str, + default=None, + help=( + "Override opcional. Default: " + "dataset/x/norm_stats.json produzido pelo split." + ), + ) + ap.add_argument("--resume", action="store_true") + ap.add_argument("--allow_missing_labels", action="store_true") + + args = ap.parse_args() + + seed_everything(args.seed) + + config_path = Path(args.config).resolve() + if not config_path.is_file(): + raise FileNotFoundError(config_path) + + with config_path.open("r", encoding="utf-8") as f: + config = json.load(f) + + train_cfg = deep_merge( + DEFAULT_TRAINING, + config.get("corridor_training", {}), + ) + + # ------------------------------------------------------------------------- + # Contrato do problema + # ------------------------------------------------------------------------- + + camera_root = resolve_project_root( + config_path, + str(config.get("camera", "oak-d")), + ) + dataset_root = camera_root / "dataset" + + resolution = config.get("resolucao", [1024, 576]) + W, H = int(resolution[0]), int(resolution[1]) + + resolution_root = ( + dataset_root + / f"{W}x{H}" + ) + + roi_inicio = float(config.get("roi_inicio", 0.0)) + roi_tamanho = float(config.get("roi_tamanho", 1.0)) + backbone = str(config.get("backbone", "nvidia/mit-b0")) + + label_classes = config.get("label_classes") + if not isinstance(label_classes, list) or len(label_classes) < 2: + raise RuntimeError( + "config['label_classes'] obrigatorio e deve conter as classes de status." + ) + + label_name_by_id = { + i: str(name) + for i, name in enumerate(label_classes) + } + num_label_classes = len(label_classes) + + labelmap_path = dataset_root / "labelmap.txt" + seg_id2label, seg_label2id = read_labelmap(labelmap_path) + + if len(seg_id2label) != 2: + raise RuntimeError( + f"Este trainer e especializado em navegavel vs nao navegavel. " + f"labelmap possui {len(seg_id2label)} classes: {seg_id2label}" + ) + + if set(seg_id2label.keys()) != {0, 1}: + raise RuntimeError( + f"Este contrato espera ids de segmentacao 0 e 1; recebido={seg_id2label}. " + "Remapeie o labelmap/masks antes do treino." + ) + + main_class_name = str(config.get("main_class_name", "navegavel")).lower() + if main_class_name not in seg_label2id: + raise RuntimeError( + f"main_class_name='{main_class_name}' nao existe no labelmap {seg_id2label}" + ) + + nav_class_id = int(seg_label2id[main_class_name]) + non_nav_ids = [cid for cid in seg_id2label if cid != nav_class_id] + if len(non_nav_ids) != 1: + raise RuntimeError("Nao foi possivel resolver classe nao navegavel.") + non_nav_class_id = int(non_nav_ids[0]) + + num_seg_classes = 2 + + # ------------------------------------------------------------------------- + # Saida + # ------------------------------------------------------------------------- + + model_family = str(config.get("modelo", "segformer_b0")) + model_name = str(config.get("model_name", "nav_mit")) + + save_dir = ( + camera_root + / "backup" + / model_family + / f"{model_name}_corridor_agri_v2" + ) + save_dir.mkdir(parents=True, exist_ok=True) + + last_path = save_dir / "last.pt" + best_operational_path = save_dir / "best_operational.pt" + best_nav_path = save_dir / "best_nav.pt" + best_status_path = save_dir / "best_status.pt" + + # Um treino novo precisa nascer sem checkpoints da rodada anterior. + if not args.resume: + for stale_path in ( + best_operational_path, + best_nav_path, + best_status_path, + last_path, + ): + if stale_path.exists(): + stale_path.unlink() + + print("[FRESH] checkpoints antigos removidos antes do novo treino.") + + # ------------------------------------------------------------------------- + # Device / performance + # ------------------------------------------------------------------------- + + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + + opt_cfg = train_cfg["optimization"] + try: + torch.set_float32_matmul_precision( + str(opt_cfg.get("matmul_precision", "high")) + ) + except Exception: + pass + + if torch.backends.cudnn.is_available(): + torch.backends.cudnn.benchmark = bool( + opt_cfg.get("cudnn_benchmark", True) + ) + + print("=" * 78) + print(f"Agri Corridor Teacher | {TRAINER_VERSION}") + print("=" * 78) + print(f"Device : {device}") + print(f"Camera root : {camera_root}") + print(f"Dataset root : {dataset_root}") + print(f"Resolution dataset : {resolution_root}") + print(f"Backbone : {backbone}") + print(f"Resolution : {W}x{H}") + print(f"ROI : begin={roi_inicio:.3f} size={roi_tamanho:.3f}") + print(f"Seg classes : {seg_id2label}") + print(f"Navegavel : id={nav_class_id} name={seg_id2label[nav_class_id]}") + print(f"Nao navegavel : id={non_nav_class_id} name={seg_id2label[non_nav_class_id]}") + print(f"Status classes : {label_name_by_id}") + print(f"Save dir : {save_dir}") + print("=" * 78) + + # ------------------------------------------------------------------------- + # Dataset + # ------------------------------------------------------------------------- + + split_train = resolution_root / "split" / "train" + split_val = resolution_root / "split" / "val" + + strict_pipeline_contract = bool( + train_cfg["data"].get( + "strict_pipeline_contract", + True, + ) + ) + + preflight_max_samples = int( + train_cfg["data"].get( + "preflight_max_samples", + 0, + ) + ) + + pipeline_audit = audit_pipeline_contract( + resolution_root=resolution_root, + train_root=split_train, + val_root=split_val, + expected_wh=(W, H), + num_seg_classes=num_seg_classes, + nav_class_id=nav_class_id, + label_classes=label_classes, + strict_contract=strict_pipeline_contract, + max_samples=preflight_max_samples, + ) + + print( + f"[PREFLIGHT] train={pipeline_audit['train']['count_total']} " + f"val={pipeline_audit['val']['count_total']} " + f"audited(train/val)=" + f"{pipeline_audit['train']['count_audited']}/" + f"{pipeline_audit['val']['count_audited']}" + ) + print( + f"[PREFLIGHT] train groups=" + f"{pipeline_audit['train']['group_counts_audited']}" + ) + print( + f"[PREFLIGHT] val groups=" + f"{pipeline_audit['val']['group_counts_audited']}" + ) + + base_train = ROISegDataset( + str(split_train), + str(save_dir), + roi_inicio, + roi_tamanho, + W, + H, + str(labelmap_path), + ) + + base_val = ROISegDataset( + str(split_val), + str(save_dir), + roi_inicio, + roi_tamanho, + W, + H, + str(labelmap_path), + ) + + strict_labels = bool(train_cfg["data"].get("strict_labels", True)) + if args.allow_missing_labels: + strict_labels = False + + ds_train = CorridorDataset( + base_train, + split_train, + label_classes=label_classes, + strict_labels=strict_labels, + ) + + ds_val = CorridorDataset( + base_val, + split_val, + label_classes=label_classes, + strict_labels=strict_labels, + ) + + print( + f"[DATA] train={len(ds_train)} val={len(ds_val)} " + f"missing_labels(train/val)=" + f"{len(ds_train.missing_labels)}/{len(ds_val.missing_labels)}" + ) + + all_group_names = sorted( + set(ds_train.group_names) + | set(ds_val.group_names) + ) + + group_name_by_id = { + i: name + for i, name in enumerate( + all_group_names + ) + } + group_id_by_name = { + name: i + for i, name in group_name_by_id.items() + } + + # Reindexa ambos datasets no mesmo espaço de IDs de grupo. + ds_train.group_id_by_name = dict( + group_id_by_name + ) + ds_train.group_name_by_id = dict( + group_name_by_id + ) + ds_train.group_ids = [ + group_id_by_name[g] + for g in ds_train.group_names + ] + + ds_val.group_id_by_name = dict( + group_id_by_name + ) + ds_val.group_name_by_id = dict( + group_name_by_id + ) + ds_val.group_ids = [ + group_id_by_name[g] + for g in ds_val.group_names + ] + + num_groups = len( + group_name_by_id + ) + + print( + f"[DATA] groups={group_name_by_id}" + ) + + # ------------------------------------------------------------------------- + # Normalizacao obrigatoria + # ------------------------------------------------------------------------- + + if args.norm_stats: + norm_stats_path_expected = Path( + args.norm_stats + ).resolve() + else: + norm_stats_path_expected = ( + resolution_root + / "norm_stats.json" + ) + + mean, std, norm_stats_path, norm_stats_raw = load_norm_stats( + norm_stats_path_expected, + device=device, + expected_resolution_wh=(W, H), + strict_contract=strict_pipeline_contract, + ) + + # O split e o norm_stats precisam descrever o dataset atual. + if strict_pipeline_contract: + train_count_stats = norm_stats_raw.get( + "train_count" + ) + val_count_stats = norm_stats_raw.get( + "val_count" + ) + + if ( + train_count_stats is not None + and int(train_count_stats) + != int( + pipeline_audit[ + "train" + ][ + "count_total" + ] + ) + ): + raise RuntimeError( + "norm_stats train_count diverge do split atual." + ) + + if ( + val_count_stats is not None + and int(val_count_stats) + != int( + pipeline_audit[ + "val" + ][ + "count_total" + ] + ) + ): + raise RuntimeError( + "norm_stats val_count diverge do split atual." + ) + + print(f"[NORM] source={norm_stats_path}") + print( + f"[NORM] scope={norm_stats_raw.get('split_scope')} " + f"images={norm_stats_raw.get('image_count')}" + ) + print( + f"[NORM] mean=" + f"{mean.detach().cpu().view(-1).tolist()}" + ) + print( + f"[NORM] std =" + f"{std.detach().cpu().view(-1).tolist()}" + ) + + # ------------------------------------------------------------------------- + # Sampler / loaders + # ------------------------------------------------------------------------- + + sampler = build_training_sampler( + ds_train, + num_label_classes, + train_cfg["sampler"], + ) + + gen = torch.Generator() + gen.manual_seed(args.seed) + + loader_kwargs = { + "batch_size": args.batch, + "num_workers": args.num_workers, + "pin_memory": device.type == "cuda", + "collate_fn": collate_corridor, + "worker_init_fn": seed_worker, + "generator": gen, + } + + if args.num_workers > 0: + loader_kwargs["persistent_workers"] = bool( + opt_cfg.get("persistent_workers", True) + ) + loader_kwargs["prefetch_factor"] = int( + opt_cfg.get("prefetch_factor", 2) + ) + + dl_train = DataLoader( + ds_train, + shuffle=(sampler is None), + sampler=sampler, + **loader_kwargs, + ) + + val_kwargs = dict(loader_kwargs) + val_kwargs["num_workers"] = max(0, args.num_workers // 2) + if val_kwargs["num_workers"] == 0: + val_kwargs.pop("persistent_workers", None) + val_kwargs.pop("prefetch_factor", None) + + dl_val = DataLoader( + ds_val, + shuffle=False, + sampler=None, + **val_kwargs, + ) + + # ------------------------------------------------------------------------- + # Modelo + # ------------------------------------------------------------------------- + + base_model = SegformerForSemanticSegmentation.from_pretrained( + backbone, + num_labels=num_seg_classes, + ignore_mismatched_sizes=True, + use_safetensors=True, + ) + + base_model.config.output_hidden_states = True + + if args.grad_ckpt: + try: + base_model.gradient_checkpointing_enable() + print("[MODEL] gradient checkpointing=ON") + except Exception as exc: + print( + f"[WARN] gradient checkpointing indisponivel: " + f"{type(exc).__name__}: {exc}" + ) + + base_model.to(device) + + with torch.no_grad(): + dummy = torch.zeros((1, 3, H, W), dtype=torch.float32, device=device) + dummy = (dummy - mean) / std + out = base_model( + pixel_values=dummy, + output_hidden_states=True, + return_dict=True, + ) + feat = out.hidden_states[-1] + feat_ch = int(feat.shape[1]) + print( + f"[MODEL] dummy seg_native={tuple(out.logits.shape)} " + f"feat_last={tuple(feat.shape)} feat_ch={feat_ch}" + ) + + sh_cfg = train_cfg["status_head"] + status_head = CorridorStatusHead( + feat_ch=feat_ch, + num_seg_classes=num_seg_classes, + num_label_classes=num_label_classes, + pool_hw=( + int(sh_cfg.get("pool_h", 3)), + int(sh_cfg.get("pool_w", 4)), + ), + hidden=int(sh_cfg.get("hidden", 256)), + dropout=float(sh_cfg.get("dropout", 0.20)), + detach_seg_summary=bool(sh_cfg.get("detach_seg_summary", True)), + seg_summary=str(sh_cfg.get("seg_summary", "probabilities")), + ).to(device) + + print(f"[MODEL] status_head={status_head.export_config()}") + + # ------------------------------------------------------------------------- + # Loss weights + # ------------------------------------------------------------------------- + + cw_cfg = train_cfg["class_weighting"] + + seg_class_weights = estimate_seg_class_weights( + ds_train, + num_classes=num_seg_classes, + ignore_index=IGNORE_INDEX, + cfg=cw_cfg, + ).to(device) + + if bool(cw_cfg.get("status_enabled", True)): + status_class_weights = estimate_status_class_weights( + ds_train, + num_label_classes=num_label_classes, + cfg=cw_cfg, + ).to(device) + else: + status_class_weights = None + + # ------------------------------------------------------------------------- + # Optimizer / scheduler / scaler + # ------------------------------------------------------------------------- + + optimizer = build_optimizer( + base_model, + status_head, + base_lr=args.lr, + default_wd=args.wd, + cfg=train_cfg["optimizer"], + ) + + updates_per_epoch = math.ceil(len(dl_train) / max(1, args.grad_accum)) + total_steps = max(1, args.epochs * updates_per_epoch) + + scheduler = build_scheduler( + optimizer, + total_steps=total_steps, + cfg=train_cfg["scheduler"], + ) + + scaler = GradScaler(enabled=bool(args.amp and device.type == "cuda")) + + augmenter = FieldCorridorAugmenter( + train_cfg["augmentation"], + ignore_index=IGNORE_INDEX, + ) + + # ------------------------------------------------------------------------- + # Resume + # ------------------------------------------------------------------------- + + start_epoch = 1 + bests = { + "operational": -1.0, + "nav_iou": -1.0, + "status_macro_f1": -1.0, + } + early = EarlyStopState() + + if args.resume and last_path.is_file(): + ckpt = load_checkpoint( + last_path, + base_model, + status_head, + optimizer=optimizer, + scheduler=scheduler, + scaler=scaler, + ) + + start_epoch = int(ckpt.get("epoch", 0)) + 1 + bests.update(ckpt.get("bests", {})) + + es = ckpt.get("early_stop", {}) or {} + early.best = float(es.get("best", bests.get("operational", -1.0))) + early.bad_epochs = int(es.get("bad_epochs", 0)) + + print( + f"[RESUME] epoch={start_epoch} " + f"bests={bests} early={early}" + ) + + resume_meta = ckpt.get("metadata", {}) or {} + gate_enabled = bool( + train_cfg.get("operational_gate", {}).get("enabled", True) + ) + if ( + gate_enabled + and resume_meta.get("operational_gate_schema") + != "agrobot.corridor.operational_gate.v1" + ): + print( + "[RESUME][OP GATE] checkpoint legado sem gate v1; " + "best_operational sera reconquistado por uma epoca elegivel." + ) + bests["operational"] = -1.0 + if best_operational_path.exists(): + best_operational_path.unlink() + + # ------------------------------------------------------------------------- + # Historico persistente de aprendizagem + # ------------------------------------------------------------------------- + + ( + metrics_jsonl_path, + metrics_json_path, + metrics_csv_path, + ) = initialize_training_history( + save_dir=save_dir, + resume=bool(args.resume), + start_epoch=start_epoch, + ) + + print(f"[METRICS] JSONL : {metrics_jsonl_path}") + print(f"[METRICS] JSON : {metrics_json_path}") + print(f"[METRICS] CSV : {metrics_csv_path}") + + # ------------------------------------------------------------------------- + # Metadata permanente para exportacao futura + # ------------------------------------------------------------------------- + + metadata_base = { + "trainer_version": TRAINER_VERSION, + "config": config, + "corridor_training": train_cfg, + "resolution_wh": [W, H], + "resolution_root": str(resolution_root), + "split_train": str(split_train), + "split_val": str(split_val), + "split_report_path": pipeline_audit.get("split_report_path"), + "split_manifest_path": pipeline_audit.get("split_manifest_path"), + "pipeline_audit": { + "train": pipeline_audit["train"], + "val": pipeline_audit["val"], + }, + "roi_inicio": roi_inicio, + "roi_tamanho": roi_tamanho, + "backbone": backbone, + "seg_id2label": seg_id2label, + "nav_class_id": nav_class_id, + "non_nav_class_id": non_nav_class_id, + "label_name_by_id": label_name_by_id, + "group_name_by_id": group_name_by_id, + "norm_stats_path": str(norm_stats_path), + "norm_channels": norm_stats_raw.get("channels"), + "norm_mean": mean.detach().cpu().view(-1).tolist(), + "norm_std": std.detach().cpu().view(-1).tolist(), + "status_head": status_head.export_config(), + "operational_gate_schema": "agrobot.corridor.operational_gate.v1", + "operational_gate_policy": copy.deepcopy( + train_cfg.get("operational_gate", {}) + ), + "onnx_target_contract": { + "input": "rgb_uint8_nhwc", + "input_color_order": "RGB", + "input_shape": [1, H, W, 3], + "outputs": { + "seg_ids": [1, H, W], + "label_probs": [1, num_label_classes], + }, + "preprocess_in_graph": True, + "argmax_in_graph": True, + }, + } + + # ------------------------------------------------------------------------- + # Train + # ------------------------------------------------------------------------- + + freeze_epochs = int( + train_cfg["optimizer"].get("freeze_encoder_epochs", 2) + ) + + early_patience = int( + train_cfg["checkpoint"].get("early_stop_patience", 18) + ) + early_delta = float( + train_cfg["checkpoint"].get("early_stop_min_delta", 3e-4) + ) + + grad_clip = float( + train_cfg["optimization"].get("grad_clip_norm", 1.0) + ) + + for epoch in range(start_epoch, args.epochs + 1): + encoder_trainable = epoch > freeze_epochs + set_encoder_trainable(base_model, encoder_trainable) + + if epoch == start_epoch or epoch == freeze_epochs + 1: + print( + f"[TRAIN] encoder_trainable={encoder_trainable} " + f"(freeze_encoder_epochs={freeze_epochs})" + ) + + print() + print("=" * 78) + print( + f"Epoch {epoch}/{args.epochs} | " + f"aug_strength={min(1.0, max(0.20, epoch / max(1, int(train_cfg['augmentation'].get('ramp_epochs', 6))))):.2f} | " + f"status_w={current_status_weight(epoch, train_cfg['loss']):.3f}" + ) + print("=" * 78) + + tr = run_epoch( + base_model=base_model, + status_head=status_head, + loader=dl_train, + device=device, + mean=mean, + std=std, + num_seg_classes=num_seg_classes, + num_label_classes=num_label_classes, + num_groups=num_groups, + group_name_by_id=group_name_by_id, + nav_class_id=nav_class_id, + non_nav_class_id=non_nav_class_id, + seg_class_weights=seg_class_weights, + status_class_weights=status_class_weights, + loss_cfg=train_cfg["loss"], + augmenter=augmenter, + optimizer=optimizer, + scheduler=scheduler, + scaler=scaler, + amp=args.amp, + train=True, + grad_accum=max(1, args.grad_accum), + grad_clip_norm=grad_clip, + epoch=epoch, + ) + + va = run_epoch( + base_model=base_model, + status_head=status_head, + loader=dl_val, + device=device, + mean=mean, + std=std, + num_seg_classes=num_seg_classes, + num_label_classes=num_label_classes, + num_groups=num_groups, + group_name_by_id=group_name_by_id, + nav_class_id=nav_class_id, + non_nav_class_id=non_nav_class_id, + seg_class_weights=seg_class_weights, + status_class_weights=status_class_weights, + loss_cfg=train_cfg["loss"], + augmenter=None, + optimizer=None, + scheduler=None, + scaler=None, + amp=args.amp_val, + train=False, + grad_accum=1, + grad_clip_norm=0.0, + epoch=epoch, + ) + + op_score = field_score( + va["seg"], + va["status"], + train_cfg["score"], + ) + + operational_gate = evaluate_operational_gate( + epoch=epoch, + val_metrics=va, + cfg=train_cfg.get("operational_gate", {}), + label_name_by_id=label_name_by_id, + ) + + print( + f"TRAIN loss={tr['loss']:.4f} " + f"seg={tr['loss_seg']:.4f} status={tr['loss_status']:.4f} " + f"navIoU={tr['seg']['nav_iou']:.4f} " + f"navF1={tr['seg']['nav_f1']:.4f} " + f"statusF1={tr['status']['macro_f1']:.4f} " + f"t={tr['time_s']:.1f}s" + ) + + print( + f"VAL loss={va['loss']:.4f} " + f"seg={va['loss_seg']:.4f} status={va['loss_status']:.4f} " + f"navIoU={va['seg']['nav_iou']:.4f} " + f"nonNavIoU={va['seg']['non_nav_iou']:.4f} " + f"navP/R/F1=" + f"{va['seg']['nav_precision']:.4f}/" + f"{va['seg']['nav_recall']:.4f}/" + f"{va['seg']['nav_f1']:.4f} " + f"unsafeNav={va['seg']['unsafe_nav_rate']:.4f} " + f"safety={va['seg']['safety']:.4f} " + f"statusAcc={va['status']['acc']:.4f} " + f"statusF1={va['status']['macro_f1']:.4f} " + f"FIELD_SCORE={op_score:.4f} " + f"t={va['time_s']:.1f}s" + ) + + print( + f"[OP GATE] {format_operational_gate(operational_gate)}" + ) + + print( + "STATUS F1 : " + + pretty_status( + va["status"]["f1_per_class"], + label_name_by_id, + ) + ) + print( + "STATUS Recall: " + + pretty_status( + va["status"]["recall_per_class"], + label_name_by_id, + ) + ) + + print("VAL POR GRUPO:") + for group_name in sorted( + va["seg_by_group"] + ): + gm = va[ + "seg_by_group" + ][group_name] + gs = va[ + "status_by_group" + ][group_name] + + print( + f" {group_name:28s} " + f"navIoU={gm['nav_iou']:.4f} " + f"nonNavIoU={gm['non_nav_iou']:.4f} " + f"unsafe={gm['unsafe_nav_rate']:.4f} " + f"statusF1={gs['macro_f1']:.4f}" + ) + + # Atualiza early stop primeiro. + if op_score > early.best + early_delta: + early.best = op_score + early.bad_epochs = 0 + else: + early.bad_epochs += 1 + + metadata_epoch = dict(metadata_base) + metadata_epoch.update({ + "epoch": epoch, + "val": va, + "train": tr, + "field_score": op_score, + "operational_gate": operational_gate, + "operational_eligible": bool(operational_gate["eligible"]), + }) + + score_improved_operational = op_score > bests["operational"] + improved_operational = ( + bool(operational_gate["eligible"]) + and score_improved_operational + ) + improved_nav = va["seg"]["nav_iou"] > bests["nav_iou"] + improved_status = va["status"]["macro_f1"] > bests["status_macro_f1"] + + # Best operacional. + if improved_operational: + bests["operational"] = float(op_score) + save_checkpoint( + best_operational_path, + base_model, + status_head, + optimizer, + scheduler, + scaler, + epoch, + bests, + early, + metadata_epoch, + ) + print( + f"[BEST OPERATIONAL] {op_score:.4f} -> " + f"{best_operational_path}" + ) + + elif score_improved_operational and not operational_gate["eligible"]: + print( + f"[OP CANDIDATE REJECTED] FIELD_SCORE={op_score:.4f} | " + f"{format_operational_gate(operational_gate)}" + ) + + # Best navegacao. + if improved_nav: + bests["nav_iou"] = float(va["seg"]["nav_iou"]) + save_checkpoint( + best_nav_path, + base_model, + status_head, + optimizer, + scheduler, + scaler, + epoch, + bests, + early, + metadata_epoch, + ) + print( + f"[BEST NAV] {bests['nav_iou']:.4f} -> " + f"{best_nav_path}" + ) + + # Best status. + if improved_status: + bests["status_macro_f1"] = float( + va["status"]["macro_f1"] + ) + save_checkpoint( + best_status_path, + base_model, + status_head, + optimizer, + scheduler, + scaler, + epoch, + bests, + early, + metadata_epoch, + ) + print( + f"[BEST STATUS] {bests['status_macro_f1']:.4f} -> " + f"{best_status_path}" + ) + + # Last sempre por ultimo para carregar os bests atualizados. + save_checkpoint( + last_path, + base_model, + status_head, + optimizer, + scheduler, + scaler, + epoch, + bests, + early, + metadata_epoch, + ) + + lrs = { + g.get("group_name", str(i)): float(g["lr"]) + for i, g in enumerate(optimizer.param_groups) + } + print(f"[LR] {lrs}") + print( + f"[ES] best={early.best:.5f} " + f"bad={early.bad_epochs}/{early_patience}" + ) + + augmentation_strength = min( + 1.0, + max( + 0.20, + epoch / max(1, int(train_cfg["augmentation"].get("ramp_epochs", 6))), + ), + ) + + gpu_memory = None + if device.type == "cuda": + gpu_memory = { + "allocated_bytes": int(torch.cuda.memory_allocated(device)), + "reserved_bytes": int(torch.cuda.memory_reserved(device)), + "max_allocated_bytes": int(torch.cuda.max_memory_allocated(device)), + "max_reserved_bytes": int(torch.cuda.max_memory_reserved(device)), + } + + epoch_history = { + "schema": "agrobot.corridor.training.metrics.v2", + "trainer_version": TRAINER_VERSION, + "created_at": time.strftime("%Y-%m-%dT%H:%M:%S"), + "epoch": int(epoch), + "epochs_requested": int(args.epochs), + "encoder_trainable": bool(encoder_trainable), + "augmentation_strength": float(augmentation_strength), + "status_weight": float(current_status_weight(epoch, train_cfg["loss"])), + "train": tr, + "val": va, + "field_score": float(op_score), + "operational_gate": operational_gate, + "operational_eligible": bool(operational_gate["eligible"]), + "operational_score_improved": bool(score_improved_operational), + "learning_rates": lrs, + "best_improved": { + "operational": bool(improved_operational), + "nav_iou": bool(improved_nav), + "status_macro_f1": bool(improved_status), + }, + "bests_after_epoch": dict(bests), + "early_stop": { + "best": float(early.best), + "bad_epochs": int(early.bad_epochs), + "patience": int(early_patience), + "min_delta": float(early_delta), + }, + "epoch_time_s": float(tr.get("time_s", 0.0) + va.get("time_s", 0.0)), + "gpu_memory": gpu_memory, + } + + save_epoch_history( + record=epoch_history, + jsonl_path=metrics_jsonl_path, + json_path=metrics_json_path, + csv_path=metrics_csv_path, + ) + + print( + f"[METRICS] epoch={epoch} persistida | " + f"FIELD_SCORE={op_score:.4f} | " + f"OP_ELIGIBLE={operational_gate['eligible']} | " + f"{metrics_csv_path.name}" + ) + + if early.bad_epochs >= early_patience: + print( + f"[EARLY STOP] FIELD_SCORE sem melhora relevante por " + f"{early.bad_epochs} epocas. " + f"(gate de elegibilidade e independente)" + ) + break + + print() + print("=" * 78) + print("Treino finalizado.") + print( + f"Best operational : {best_operational_path} " + f"({'OK' if best_operational_path.is_file() else 'NAO GERADO - nenhum elegivel'})" + ) + print(f"Best nav : {best_nav_path}") + print(f"Best status : {best_status_path}") + print(f"Metrics JSONL : {metrics_jsonl_path}") + print(f"Metrics JSON : {metrics_json_path}") + print(f"Metrics CSV : {metrics_csv_path}") + print("=" * 78) + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/oak-d/_5_test_corridor.py b/Python/OAK/datasets/oak-d/_5_test_corridor.py new file mode 100644 index 000000000..8f4e11478 --- /dev/null +++ b/Python/OAK/datasets/oak-d/_5_test_corridor.py @@ -0,0 +1,4052 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +test_corridor_agri_v2.py +======================== + +Viewer/teste manual do modelo frontal OAK-D Lite, alinhado ao pipeline 2026-09. + +Compatível com checkpoints do: + train_corridor_agri_v2.py + +Tarefa: + 1) segmentação binária: + - não navegável + - navegável + + 2) classificação global do estado do corredor: + - config["label_classes"] + +Fontes suportadas: + A) split com Ground Truth: + dataset/x/split/train + dataset/x/split/val + ou qualquer pasta com: + images/ + masks/ opcional + labels/ opcional + + B) pasta simples: + qualquer pasta contendo PNG/JPG/etc. + Sem GT é tratada como preview/imagens externas. + + C) câmera: + OAK-D Lite em tempo real. + Captura 1080p por padrão para manter paridade com a coleta/produção, + e aplica o MESMO preprocess usado para pasta/dataset. + +Exemplos: + # Val padrão do pipeline atual: + python test_corridor_agri_v2.py + + # Train: + python test_corridor_agri_v2.py --split train + + # Pasta estruturada: + python test_corridor_agri_v2.py --input dataset/1024x576/split/val + + # Pasta só com imagens: + python test_corridor_agri_v2.py --input C:/meus_previews + + # Câmera: + python test_corridor_agri_v2.py --camera + + # Avaliar todos sem janela: + python test_corridor_agri_v2.py --split val --no-view --save-dir outputs_val + +Teclas: + D / seta direita / SPACE = próxima + A / seta esquerda = anterior + S = salva painel atual + Q / ESC = sair + +Decisões de projeto: + - checkpoint é a fonte de verdade de arquitetura, resolução, status head + e normalização; + - não usa o LabelHead antigo; + - não usa input quadrado 512; + - GT é opcional; + - camera/pasta/dataset passam pelo mesmo preprocess; + - em dataset com GT mostra erro perigoso (falso navegável) separadamente. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import os +import re +import time +from collections import Counter +from dataclasses import dataclass +from pathlib import Path +from typing import Dict, Iterable, List, Optional, Sequence, Tuple + +import cv2 +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from transformers import SegformerForSemanticSegmentation + + +# ============================================================================= +# Constants +# ============================================================================= + +TESTER_VERSION = "agri_corridor_manual_test_v2.0" +IGNORE_INDEX = 255 + +IMG_EXTS = { + ".png", + ".jpg", + ".jpeg", + ".bmp", + ".webp", + ".tif", + ".tiff", +} + +DEFAULT_OVERLAY_ALPHA = 0.45 + + +# ============================================================================= +# Data +# ============================================================================= + +@dataclass(frozen=True) +class SegClass: + cid: int + name: str + rgb: Tuple[int, int, int] + + @property + def bgr(self) -> Tuple[int, int, int]: + r, g, b = self.rgb + return b, g, r + + +@dataclass(frozen=True) +class Sample: + base: str + image: Path + mask: Optional[Path] + label: Optional[Path] + group: str + source_root: Path + + +@dataclass +class RuntimeContract: + checkpoint_path: Path + checkpoint_epoch: int + backbone: str + resolution_wh: Tuple[int, int] + roi_inicio: float + roi_tamanho: float + + seg_id2label: Dict[int, str] + nav_class_id: int + non_nav_class_id: int + + label_name_by_id: Dict[int, str] + group_name_by_id: Dict[int, str] + + norm_mean: np.ndarray + norm_std: np.ndarray + norm_channels: List[str] + + status_head_cfg: dict + trainer_version: str + + +# ============================================================================= +# Generic helpers +# ============================================================================= + +def natural_key(text: str): + parts = re.split(r"(\d+)", str(text)) + return [ + int(p) if p.isdigit() else p.lower() + for p in parts + ] + + +def safe_json_load(path: Path) -> dict: + with path.open("r", encoding="utf-8") as f: + data = json.load(f) + + if not isinstance(data, dict): + raise ValueError( + f"Esperava objeto JSON em {path}" + ) + + return data + + +def find_project_root(config_path: Path) -> Path: + root = config_path.resolve().parent + + if not (root / "dataset").is_dir(): + raise FileNotFoundError( + f"dataset/ não encontrado ao lado de {config_path}. " + "Execute este script dentro da pasta oak-d ou passe --config correto." + ) + + return root + + +def map_by_stem( + folder: Path, + exts: set[str], +) -> Dict[str, Path]: + if not folder.is_dir(): + return {} + + result: Dict[str, Path] = {} + + for p in sorted( + folder.iterdir(), + key=lambda x: natural_key(x.name), + ): + if ( + not p.is_file() + or p.suffix.lower() not in exts + ): + continue + + if p.stem in result: + raise RuntimeError( + f"Stem duplicado {p.stem!r} em {folder}: " + f"{result[p.stem].name} / {p.name}" + ) + + result[p.stem] = p + + return result + + +def ensure_bgr(path: Path) -> np.ndarray: + img = cv2.imread( + str(path), + cv2.IMREAD_COLOR, + ) + + if img is None: + raise RuntimeError( + f"Falha ao ler imagem: {path}" + ) + + return img + + +def clamp01(x: float) -> float: + return max( + 0.0, + min( + 1.0, + float(x), + ), + ) + + +# ============================================================================= +# Labelmap +# ============================================================================= + +def _parse_rgb_triplet(text: str) -> Optional[Tuple[int, int, int]]: + nums = re.findall( + r"-?\d+", + str(text), + ) + + if len(nums) < 3: + return None + + rgb = tuple( + int(x) + for x in nums[:3] + ) + + if not all( + 0 <= x <= 255 + for x in rgb + ): + return None + + return rgb # type: ignore + + +def load_labelmap( + path: Path, +) -> Dict[int, SegClass]: + """ + Formatos aceitos: + navegavel: 40,220,80 + 1 navegavel 40 220 80 + 1:navegavel:40,220,80 + """ + if not path.is_file(): + raise FileNotFoundError( + f"labelmap não encontrado: {path}" + ) + + classes: Dict[int, SegClass] = {} + next_id = 0 + + for raw in path.read_text( + encoding="utf-8" + ).splitlines(): + line = raw.strip() + + if ( + not line + or line.startswith("#") + ): + continue + + lower = line.lower() + + if any( + token in lower + for token in ( + "ignore", + "void", + "background_ignore", + ) + ): + continue + + cid: Optional[int] = None + name: Optional[str] = None + rgb: Optional[ + Tuple[int, int, int] + ] = None + + # "name: R,G,B" ou "id:name:R,G,B" + colon = [ + p.strip() + for p in line.split(":") + ] + + if len(colon) >= 2: + try: + possible_id = int( + colon[0] + ) + + if len(colon) >= 3: + cid = possible_id + name = colon[1] + rgb = _parse_rgb_triplet( + ":".join(colon[2:]) + ) + except ValueError: + name = colon[0] + rgb = _parse_rgb_triplet( + ":".join(colon[1:]) + ) + + # "id name R G B" + if ( + name is None + or rgb is None + ): + tokens = re.split( + r"[\s,;]+", + line, + ) + tokens = [ + x + for x in tokens + if x + ] + + if ( + len(tokens) >= 5 + and tokens[0].lstrip("-").isdigit() + ): + cid = int(tokens[0]) + name = tokens[1] + + try: + rgb = ( + int(tokens[2]), + int(tokens[3]), + int(tokens[4]), + ) + except Exception: + rgb = None + + if ( + name is None + or rgb is None + ): + raise ValueError( + f"Linha de labelmap não reconhecida: {line!r}" + ) + + if cid is None: + while next_id in classes: + next_id += 1 + cid = next_id + + classes[int(cid)] = SegClass( + cid=int(cid), + name=str(name), + rgb=( + int(rgb[0]), + int(rgb[1]), + int(rgb[2]), + ), + ) + + next_id = max( + next_id, + int(cid) + 1, + ) + + if len(classes) != 2: + raise RuntimeError( + f"Tester de corredor espera exatamente 2 classes; " + f"labelmap={classes}" + ) + + if set(classes) != {0, 1}: + raise RuntimeError( + f"Tester espera IDs 0/1; recebido={sorted(classes)}" + ) + + return classes + + +def colorize_ids( + mask_ids: np.ndarray, + classes: Dict[int, SegClass], +) -> np.ndarray: + out = np.zeros( + ( + mask_ids.shape[0], + mask_ids.shape[1], + 3, + ), + dtype=np.uint8, + ) + + for cid, cls in classes.items(): + out[ + mask_ids == int(cid) + ] = np.asarray( + cls.bgr, + dtype=np.uint8, + ) + + out[ + mask_ids == IGNORE_INDEX + ] = np.asarray( + (180, 0, 180), + dtype=np.uint8, + ) + + return out + + +def decode_color_mask( + raw_bgr: np.ndarray, + classes: Dict[int, SegClass], +) -> np.ndarray: + ids = np.full( + raw_bgr.shape[:2], + IGNORE_INDEX, + dtype=np.uint8, + ) + + for cid, cls in classes.items(): + bgr = np.asarray( + cls.bgr, + dtype=np.uint8, + ) + + match = np.all( + raw_bgr[:, :, :3] + == bgr[None, None, :], + axis=2, + ) + + ids[match] = int(cid) + + return ids + + +def load_gt_mask( + path: Optional[Path], + classes: Dict[int, SegClass], +) -> Optional[np.ndarray]: + if path is None: + return None + + raw = cv2.imread( + str(path), + cv2.IMREAD_UNCHANGED, + ) + + if raw is None: + raise RuntimeError( + f"Falha ao ler GT mask: {path}" + ) + + if raw.ndim == 2: + return raw.astype( + np.uint8, + copy=False, + ) + + if ( + raw.ndim == 3 + and raw.shape[2] >= 3 + ): + return decode_color_mask( + raw[:, :, :3], + classes, + ) + + raise RuntimeError( + f"Formato de GT mask inválido: {path} shape={raw.shape}" + ) + + +# ============================================================================= +# Discovery +# ============================================================================= + +def discover_flat_dataset( + root: Path, +) -> List[Sample]: + images = map_by_stem( + root / "images", + IMG_EXTS, + ) + + if not images: + return [] + + masks = map_by_stem( + root / "masks", + {".png", ".jpg", ".jpeg", ".bmp", ".tif", ".tiff"}, + ) + + labels = map_by_stem( + root / "labels", + {".json"}, + ) + + result: List[Sample] = [] + + for base, image in sorted( + images.items(), + key=lambda kv: natural_key(kv[0]), + ): + group = "unknown" + + label = labels.get(base) + + if label is not None: + try: + data = safe_json_load( + label + ) + group = str( + data.get( + "split_source_group", + data.get( + "group", + "unknown", + ), + ) + ) + except Exception: + group = "unknown" + + result.append( + Sample( + base=base, + image=image, + mask=masks.get(base), + label=label, + group=group, + source_root=root, + ) + ) + + return result + + +def discover_grouped_dataset( + root: Path, +) -> List[Sample]: + """ + Aceita: + root/group//{images,masks,labels} + ou: + root//{images,masks,labels} + """ + group_root = ( + root / "group" + if (root / "group").is_dir() + else root + ) + + samples: List[Sample] = [] + + image_dirs = sorted( + { + p + for p in group_root.rglob("images") + if p.is_dir() + }, + key=lambda p: natural_key( + str(p) + ), + ) + + for images_dir in image_dirs: + base_dir = images_dir.parent + + try: + group = str( + base_dir.relative_to( + group_root + ) + ).replace( + "\\", + "/", + ) + except Exception: + group = base_dir.name + + images = map_by_stem( + images_dir, + IMG_EXTS, + ) + masks = map_by_stem( + base_dir / "masks", + {".png", ".jpg", ".jpeg", ".bmp", ".tif", ".tiff"}, + ) + labels = map_by_stem( + base_dir / "labels", + {".json"}, + ) + + for stem, image in sorted( + images.items(), + key=lambda kv: natural_key(kv[0]), + ): + samples.append( + Sample( + base=stem, + image=image, + mask=masks.get(stem), + label=labels.get(stem), + group=group, + source_root=root, + ) + ) + + return samples + + +def discover_plain_folder( + root: Path, +) -> List[Sample]: + result: List[Sample] = [] + + if not root.is_dir(): + return result + + for p in sorted( + root.iterdir(), + key=lambda x: natural_key(x.name), + ): + if ( + p.is_file() + and p.suffix.lower() + in IMG_EXTS + ): + result.append( + Sample( + base=p.stem, + image=p, + mask=None, + label=None, + group="external", + source_root=root, + ) + ) + + return result + + +def discover_samples( + root: Path, +) -> Tuple[List[Sample], str]: + if not root.is_dir(): + raise FileNotFoundError( + f"Entrada não encontrada: {root}" + ) + + # 1) split novo / pasta estruturada plana. + flat = discover_flat_dataset( + root + ) + + if flat: + return flat, "dataset_flat" + + # 2) original/group ou normalized/group. + grouped = discover_grouped_dataset( + root + ) + + if grouped: + return grouped, "dataset_grouped" + + # 3) pasta de imagens. + plain = discover_plain_folder( + root + ) + + if plain: + return plain, "image_folder" + + raise RuntimeError( + f"Nenhuma imagem reconhecida em: {root}" + ) + + +def read_gt_label( + sample: Sample, + label_name_by_id: Dict[int, str], +) -> Tuple[Optional[int], Optional[str]]: + if sample.label is None: + return None, None + + data = safe_json_load( + sample.label + ) + + lid = data.get( + "label_id" + ) + + name = ( + data.get("estado_corredor") + or data.get("label") + or data.get("state") + ) + + if lid is not None: + lid = int(lid) + + if name is None: + name = label_name_by_id.get( + lid + ) + + return ( + int(lid) + if lid is not None + else None, + str(name) + if name is not None + else None, + ) + + +# ============================================================================= +# Status head exact contract +# ============================================================================= + +class CorridorStatusHead(nn.Module): + """ + Espelho da cabeça usada por train_corridor_agri_v2.py. + + Não é o LabelHead legado: + - feature profunda permanece nativa; + - AdaptiveAvgPool 3x4 (ou metadata); + - resumo da segmentação separado; + - concatena vetores compactos. + """ + + def __init__( + self, + feat_ch: int, + num_seg_classes: int, + num_label_classes: int, + pool_hw: Tuple[int, int], + hidden: int, + dropout: float, + detach_seg_summary: bool, + seg_summary: str, + ): + super().__init__() + + self.feat_ch = int( + feat_ch + ) + self.num_seg_classes = int( + num_seg_classes + ) + self.num_label_classes = int( + num_label_classes + ) + self.pool_hw = ( + int(pool_hw[0]), + int(pool_hw[1]), + ) + self.detach_seg_summary = bool( + detach_seg_summary + ) + self.seg_summary = str( + seg_summary + ).lower() + + pooled_cells = ( + self.pool_hw[0] + * self.pool_hw[1] + ) + + in_dim = ( + self.feat_ch + + self.num_seg_classes + ) * pooled_cells + + self.feat_pool = nn.AdaptiveAvgPool2d( + self.pool_hw + ) + self.seg_pool = nn.AdaptiveAvgPool2d( + self.pool_hw + ) + + self.net = nn.Sequential( + nn.Linear( + in_dim, + int(hidden), + ), + nn.ReLU( + inplace=True + ), + nn.Dropout( + float(dropout) + ), + nn.Linear( + int(hidden), + self.num_label_classes, + ), + ) + + def forward( + self, + feat: torch.Tensor, + seg_logits_native: torch.Tensor, + ) -> torch.Tensor: + feat_grid = self.feat_pool( + feat + ).flatten(1) + + seg_src = ( + seg_logits_native.detach() + if self.detach_seg_summary + else seg_logits_native + ) + + if ( + self.seg_summary + == "probabilities" + ): + seg_src = torch.softmax( + seg_src, + dim=1, + ) + elif ( + self.seg_summary + != "logits" + ): + raise RuntimeError( + f"seg_summary inválido: " + f"{self.seg_summary}" + ) + + seg_grid = self.seg_pool( + seg_src + ).flatten(1) + + x = torch.cat( + [ + feat_grid, + seg_grid, + ], + dim=1, + ) + + return self.net( + x + ) + + +# ============================================================================= +# Checkpoint +# ============================================================================= + +def default_checkpoint_path( + project_root: Path, + config: dict, + which: str, +) -> Path: + model_family = str( + config.get( + "modelo", + "segformer_b0", + ) + ) + model_name = str( + config.get( + "model_name", + "nav_mit", + ) + ) + + save_dir = ( + project_root + / "backup" + / model_family + / f"{model_name}_corridor_agri_v2" + ) + + filename = { + "operational": "best_operational.pt", + "nav": "best_nav.pt", + "status": "best_status.pt", + "last": "last.pt", + }[which] + + return save_dir / filename + + +def _int_key_dict( + value, +) -> Dict[int, str]: + if not isinstance( + value, + dict, + ): + return {} + + return { + int(k): str(v) + for k, v in value.items() + } + + +def build_runtime_from_checkpoint( + checkpoint_path: Path, + config: dict, + device: torch.device, +): + if not checkpoint_path.is_file(): + raise FileNotFoundError( + f"Checkpoint não encontrado: {checkpoint_path}" + ) + + ckpt = torch.load( + checkpoint_path, + map_location="cpu", + weights_only=False, + ) + + if "model" not in ckpt: + raise RuntimeError( + "Checkpoint não possui state_dict 'model'." + ) + + if "status_head" not in ckpt: + raise RuntimeError( + "Checkpoint não possui 'status_head'. " + "Este tester é para train_corridor_agri_v2.py, " + "não para o LabelHead antigo." + ) + + metadata = ckpt.get( + "metadata", + {}, + ) or {} + + trainer_version = str( + ckpt.get( + "trainer_version", + metadata.get( + "trainer_version", + "unknown", + ), + ) + ) + + resolution = metadata.get( + "resolution_wh", + config.get( + "resolucao", + [1024, 576], + ), + ) + + W = int( + resolution[0] + ) + H = int( + resolution[1] + ) + + backbone = str( + metadata.get( + "backbone", + config.get( + "backbone", + "nvidia/mit-b0", + ), + ) + ) + + seg_id2label = _int_key_dict( + metadata.get( + "seg_id2label", + {}, + ) + ) + + if not seg_id2label: + seg_id2label = { + 0: "naonavegavel", + 1: "navegavel", + } + + nav_class_id = int( + metadata.get( + "nav_class_id", + 1, + ) + ) + non_nav_class_id = int( + metadata.get( + "non_nav_class_id", + 1 - nav_class_id, + ) + ) + + label_name_by_id = _int_key_dict( + metadata.get( + "label_name_by_id", + {}, + ) + ) + + if not label_name_by_id: + raw = config.get( + "label_classes", + [], + ) + label_name_by_id = { + i: str(name) + for i, name in enumerate( + raw + ) + } + + if not label_name_by_id: + raise RuntimeError( + "Não consegui resolver classes da segunda cabeça." + ) + + group_name_by_id = _int_key_dict( + metadata.get( + "group_name_by_id", + {}, + ) + ) + + mean = metadata.get( + "norm_mean" + ) + std = metadata.get( + "norm_std" + ) + channels = metadata.get( + "norm_channels", + ["R", "G", "B"], + ) + + if ( + mean is None + or std is None + ): + raise RuntimeError( + "Checkpoint não possui norm_mean/norm_std em metadata. " + "O tester v2 exige o contrato do trainer v2." + ) + + mean_np = np.asarray( + mean, + dtype=np.float32, + ).reshape(3) + std_np = np.asarray( + std, + dtype=np.float32, + ).reshape(3) + + if not np.all( + np.isfinite(mean_np) + ) or not np.all( + np.isfinite(std_np) + ): + raise RuntimeError( + "Checkpoint possui norm stats NaN/Inf." + ) + + if np.any( + std_np <= 0 + ): + raise RuntimeError( + f"Checkpoint possui std inválido: {std_np.tolist()}" + ) + + status_sd = ckpt[ + "status_head" + ] + + status_meta = dict( + metadata.get( + "status_head", + {}, + ) or {} + ) + + # Inferências seguras diretamente do state_dict. + if "net.0.weight" not in status_sd: + raise RuntimeError( + "status_head state_dict não possui net.0.weight." + ) + + hidden = int( + status_sd[ + "net.0.weight" + ].shape[0] + ) + + final_weight_key = None + + for candidate in ( + "net.3.weight", + "net.4.weight", + ): + if candidate in status_sd: + final_weight_key = candidate + break + + if final_weight_key is None: + raise RuntimeError( + "Não consegui inferir saída do status_head." + ) + + num_label_classes = int( + status_sd[ + final_weight_key + ].shape[0] + ) + + feat_ch = int( + status_meta.get( + "feat_ch", + 0, + ) + ) + + num_seg_classes = int( + status_meta.get( + "num_seg_classes", + len(seg_id2label), + ) + ) + + pool_hw_raw = status_meta.get( + "pool_hw", + [3, 4], + ) + pool_hw = ( + int(pool_hw_raw[0]), + int(pool_hw_raw[1]), + ) + + detach_seg_summary = bool( + status_meta.get( + "detach_seg_summary", + True, + ) + ) + seg_summary = str( + status_meta.get( + "seg_summary", + "probabilities", + ) + ) + + base_model = ( + SegformerForSemanticSegmentation + .from_pretrained( + backbone, + num_labels=num_seg_classes, + ignore_mismatched_sizes=True, + use_safetensors=True, + ) + ) + + base_model.config.output_hidden_states = True + + base_model.load_state_dict( + ckpt["model"], + strict=True, + ) + + base_model.to( + device + ).eval() + + # Metadata antiga pode não ter feat_ch. Descobre com dummy. + if feat_ch <= 0: + with torch.inference_mode(): + dummy = torch.zeros( + ( + 1, + 3, + H, + W, + ), + dtype=torch.float32, + device=device, + ) + + mean_t = torch.tensor( + mean_np, + dtype=torch.float32, + device=device, + ).view( + 1, + 3, + 1, + 1, + ) + std_t = torch.tensor( + std_np, + dtype=torch.float32, + device=device, + ).view( + 1, + 3, + 1, + 1, + ) + + out = base_model( + pixel_values=( + dummy + - mean_t + ) / std_t, + output_hidden_states=True, + return_dict=True, + ) + + feat_ch = int( + out.hidden_states[-1].shape[1] + ) + + status_head = CorridorStatusHead( + feat_ch=feat_ch, + num_seg_classes=num_seg_classes, + num_label_classes=num_label_classes, + pool_hw=pool_hw, + hidden=hidden, + dropout=0.0, # irrelevante em eval + detach_seg_summary=detach_seg_summary, + seg_summary=seg_summary, + ) + + status_head.load_state_dict( + status_sd, + strict=True, + ) + + status_head.to( + device + ).eval() + + contract = RuntimeContract( + checkpoint_path=checkpoint_path, + checkpoint_epoch=int( + ckpt.get( + "epoch", + -1, + ) + ), + backbone=backbone, + resolution_wh=( + W, + H, + ), + roi_inicio=float( + metadata.get( + "roi_inicio", + config.get( + "roi_inicio", + 0.0, + ), + ) + ), + roi_tamanho=float( + metadata.get( + "roi_tamanho", + config.get( + "roi_tamanho", + 1.0, + ), + ) + ), + seg_id2label=seg_id2label, + nav_class_id=nav_class_id, + non_nav_class_id=non_nav_class_id, + label_name_by_id=label_name_by_id, + group_name_by_id=group_name_by_id, + norm_mean=mean_np, + norm_std=std_np, + norm_channels=[ + str(x) + for x in channels + ], + status_head_cfg={ + "feat_ch": feat_ch, + "num_seg_classes": num_seg_classes, + "num_label_classes": num_label_classes, + "pool_hw": list( + pool_hw + ), + "hidden": hidden, + "detach_seg_summary": detach_seg_summary, + "seg_summary": seg_summary, + }, + trainer_version=trainer_version, + ) + + return ( + base_model, + status_head, + contract, + ckpt, + ) + + +# ============================================================================= +# Preprocess / inference +# ============================================================================= + +def compute_roi( + h: int, + roi_inicio: float, + roi_tamanho: float, +) -> Tuple[int, int]: + y0 = int( + round( + h + * float( + roi_inicio + ) + ) + ) + y1 = int( + round( + h + * min( + 1.0, + float( + roi_inicio + + roi_tamanho + ), + ) + ) + ) + + y0 = max( + 0, + min( + h - 1, + y0, + ), + ) + y1 = max( + y0 + 1, + min( + h, + y1, + ), + ) + + return y0, y1 + + +def prepare_input( + img_bgr: np.ndarray, + contract: RuntimeContract, + device: torch.device, +): + h0, w0 = img_bgr.shape[:2] + + y0, y1 = compute_roi( + h0, + contract.roi_inicio, + contract.roi_tamanho, + ) + + roi_bgr = img_bgr[ + y0:y1, + 0:w0, + ] + + W, H = contract.resolution_wh + + resized_bgr = cv2.resize( + roi_bgr, + ( + W, + H, + ), + interpolation=cv2.INTER_AREA, + ) + + rgb = cv2.cvtColor( + resized_bgr, + cv2.COLOR_BGR2RGB, + ) + + arr = ( + rgb.astype( + np.float32 + ) + / 255.0 + ) + + arr = np.transpose( + arr, + ( + 2, + 0, + 1, + ), + ) + + mean = contract.norm_mean.reshape( + 3, + 1, + 1, + ) + std = contract.norm_std.reshape( + 3, + 1, + 1, + ) + + arr = ( + arr + - mean + ) / np.maximum( + std, + 1e-6, + ) + + x = torch.from_numpy( + np.ascontiguousarray( + arr + ) + ).unsqueeze(0).to( + device, + non_blocking=True, + ) + + return ( + x, + { + "y0": y0, + "y1": y1, + "source_wh": ( + w0, + h0, + ), + "roi_wh": ( + w0, + y1 - y0, + ), + }, + ) + + +@torch.inference_mode() +def infer( + base_model, + status_head, + img_bgr: np.ndarray, + contract: RuntimeContract, + device: torch.device, +) -> dict: + x, geometry = prepare_input( + img_bgr, + contract, + device, + ) + + if device.type == "cuda": + torch.cuda.synchronize() + + t0 = time.perf_counter() + + out = base_model( + pixel_values=x, + output_hidden_states=True, + return_dict=True, + ) + + seg_logits_native = out.logits + feat = out.hidden_states[-1] + + status_logits = status_head( + feat, + seg_logits_native, + ) + + W, H = contract.resolution_wh + + seg_logits_model = seg_logits_native + + if ( + seg_logits_model.shape[-2:] + != ( + H, + W, + ) + ): + seg_logits_model = F.interpolate( + seg_logits_model, + size=( + H, + W, + ), + mode="bilinear", + align_corners=False, + ) + + pred_model = torch.argmax( + seg_logits_model, + dim=1, + )[0] + + status_probs = torch.softmax( + status_logits, + dim=1, + )[0] + + if device.type == "cuda": + torch.cuda.synchronize() + + infer_ms = ( + time.perf_counter() + - t0 + ) * 1000.0 + + pred_model_np = ( + pred_model + .detach() + .cpu() + .numpy() + .astype( + np.uint8 + ) + ) + + probs_np = ( + status_probs + .detach() + .cpu() + .numpy() + .astype( + np.float32 + ) + ) + + label_id = int( + np.argmax( + probs_np + ) + ) + + # Volta a segmentação para a geometria original para visualização. + w0, h0 = geometry[ + "source_wh" + ] + y0 = geometry["y0"] + y1 = geometry["y1"] + + pred_roi_source = cv2.resize( + pred_model_np, + ( + w0, + y1 - y0, + ), + interpolation=cv2.INTER_NEAREST, + ) + + pred_full = np.full( + ( + h0, + w0, + ), + IGNORE_INDEX, + dtype=np.uint8, + ) + + pred_full[ + y0:y1, + 0:w0, + ] = pred_roi_source + + return { + "pred_model": pred_model_np, + "pred_full": pred_full, + "status_probs": probs_np, + "label_id": label_id, + "label_conf": float( + probs_np[ + label_id + ] + ), + "infer_ms": float( + infer_ms + ), + "geometry": geometry, + } + + +def gt_to_model_space( + gt_full: Optional[np.ndarray], + geometry: dict, + model_wh: Tuple[int, int], +) -> Optional[np.ndarray]: + if gt_full is None: + return None + + y0 = int( + geometry["y0"] + ) + y1 = int( + geometry["y1"] + ) + + roi = gt_full[ + y0:y1, + :, + ] + + W, H = model_wh + + return cv2.resize( + roi, + ( + W, + H, + ), + interpolation=cv2.INTER_NEAREST, + ) + + +# ============================================================================= +# Metrics +# ============================================================================= + +def confusion_from_arrays( + pred: np.ndarray, + gt: np.ndarray, + num_classes: int, +) -> np.ndarray: + valid = ( + gt != IGNORE_INDEX + ) + + p = pred[ + valid + ].astype( + np.int64 + ) + t = gt[ + valid + ].astype( + np.int64 + ) + + valid2 = ( + (t >= 0) + & (t < num_classes) + & (p >= 0) + & (p < num_classes) + ) + + p = p[ + valid2 + ] + t = t[ + valid2 + ] + + cm = np.zeros( + ( + num_classes, + num_classes, + ), + dtype=np.int64, + ) + + if t.size: + idx = ( + t + * num_classes + + p + ) + + bins = np.bincount( + idx, + minlength=( + num_classes + * num_classes + ), + ) + + cm += bins.reshape( + num_classes, + num_classes, + ) + + return cm + + +def metrics_from_cm( + cm: np.ndarray, + nav_class_id: int, + non_nav_class_id: int, +) -> dict: + cm = cm.astype( + np.float64 + ) + + tp = np.diag( + cm + ) + fp = cm.sum( + axis=0 + ) - tp + fn = cm.sum( + axis=1 + ) - tp + + iou = tp / np.maximum( + tp + fp + fn, + 1e-12, + ) + + precision = tp / np.maximum( + tp + fp, + 1e-12, + ) + recall = tp / np.maximum( + tp + fn, + 1e-12, + ) + + f1 = ( + 2 + * precision + * recall + / np.maximum( + precision + recall, + 1e-12, + ) + ) + + acc = ( + tp.sum() + / max( + 1.0, + cm.sum(), + ) + ) + + gt_non_nav = cm[ + non_nav_class_id, + :, + ].sum() + + unsafe = ( + cm[ + non_nav_class_id, + nav_class_id, + ] + / max( + 1.0, + gt_non_nav, + ) + ) + + return { + "acc": float( + acc + ), + "miou": float( + iou.mean() + ), + "iou_per_class": iou.tolist(), + "f1_per_class": f1.tolist(), + "nav_iou": float( + iou[ + nav_class_id + ] + ), + "nav_f1": float( + f1[ + nav_class_id + ] + ), + "non_nav_iou": float( + iou[ + non_nav_class_id + ] + ), + "unsafe_nav_rate": float( + unsafe + ), + "safety": float( + 1.0 - unsafe + ), + } + + +def status_metrics_from_cm( + cm: np.ndarray, +) -> dict: + cm = cm.astype( + np.float64 + ) + + tp = np.diag( + cm + ) + fp = cm.sum( + axis=0 + ) - tp + fn = cm.sum( + axis=1 + ) - tp + support = cm.sum( + axis=1 + ) + + precision = tp / np.maximum( + tp + fp, + 1e-12, + ) + recall = tp / np.maximum( + tp + fn, + 1e-12, + ) + f1 = ( + 2 + * precision + * recall + / np.maximum( + precision + recall, + 1e-12, + ) + ) + + present = ( + support > 0 + ) + + macro_f1 = ( + float( + f1[ + present + ].mean() + ) + if np.any( + present + ) + else 0.0 + ) + + acc = ( + float( + tp.sum() + / max( + 1.0, + cm.sum(), + ) + ) + ) + + return { + "acc": acc, + "macro_f1": macro_f1, + "f1_per_class": f1.tolist(), + "recall_per_class": recall.tolist(), + "support_per_class": support.tolist(), + } + + +# ============================================================================= +# Visualization +# ============================================================================= + +def overlay_ids( + img_bgr: np.ndarray, + mask_ids: np.ndarray, + classes: Dict[int, SegClass], + alpha: float, +) -> np.ndarray: + colored = colorize_ids( + mask_ids, + classes, + ) + + valid = ( + mask_ids + != IGNORE_INDEX + ) + + out = img_bgr.copy() + + if np.any( + valid + ): + blended = cv2.addWeighted( + img_bgr, + 1.0 - alpha, + colored, + alpha, + 0.0, + ) + + out[ + valid + ] = blended[ + valid + ] + + return out + + +def error_visual( + img_bgr: np.ndarray, + gt: Optional[np.ndarray], + pred: np.ndarray, + nav_class_id: int, + non_nav_class_id: int, +) -> np.ndarray: + """ + Vermelho = falso navegável (perigoso) + Amarelo = navegável bloqueado + Verde = acerto de classe navegável + Cinza = demais acertos + Magenta = GT ignore/desconhecido + """ + out = ( + img_bgr.astype( + np.float32 + ) + * 0.45 + ).astype( + np.uint8 + ) + + if gt is None: + return out + + if pred.shape != gt.shape: + pred = cv2.resize( + pred, + ( + gt.shape[1], + gt.shape[0], + ), + interpolation=cv2.INTER_NEAREST, + ) + + valid = ( + gt != IGNORE_INDEX + ) + + unsafe = ( + valid + & ( + gt + == non_nav_class_id + ) + & ( + pred + == nav_class_id + ) + ) + + blocked = ( + valid + & ( + gt + == nav_class_id + ) + & ( + pred + == non_nav_class_id + ) + ) + + correct_nav = ( + valid + & ( + gt + == nav_class_id + ) + & ( + pred + == nav_class_id + ) + ) + + correct_other = ( + valid + & ( + gt + == pred + ) + & ~correct_nav + ) + + out[ + correct_other + ] = ( + 90, + 90, + 90, + ) + out[ + correct_nav + ] = ( + 60, + 155, + 60, + ) + out[ + blocked + ] = ( + 0, + 210, + 255, + ) + out[ + unsafe + ] = ( + 0, + 0, + 255, + ) + out[ + ~valid + ] = ( + 180, + 0, + 180, + ) + + return out + + +def put_title( + image: np.ndarray, + title: str, +) -> np.ndarray: + out = image.copy() + + overlay = out.copy() + cv2.rectangle( + overlay, + ( + 0, + 0, + ), + ( + out.shape[1], + 42, + ), + ( + 15, + 15, + 15, + ), + -1, + ) + + out = cv2.addWeighted( + overlay, + 0.78, + out, + 0.22, + 0.0, + ) + + cv2.putText( + out, + str(title), + ( + 12, + 29, + ), + cv2.FONT_HERSHEY_SIMPLEX, + 0.72, + ( + 245, + 245, + 245, + ), + 2, + cv2.LINE_AA, + ) + + return out + + +def resize_letterbox( + image: np.ndarray, + width: int, + height: int, +) -> np.ndarray: + h, w = image.shape[:2] + + scale = min( + width / max( + 1, + w, + ), + height / max( + 1, + h, + ), + ) + + nw = max( + 1, + int( + round( + w + * scale + ) + ), + ) + nh = max( + 1, + int( + round( + h + * scale + ) + ), + ) + + resized = cv2.resize( + image, + ( + nw, + nh, + ), + interpolation=( + cv2.INTER_AREA + if scale < 1.0 + else cv2.INTER_LINEAR + ), + ) + + canvas = np.zeros( + ( + height, + width, + 3, + ), + dtype=np.uint8, + ) + canvas[:] = ( + 24, + 24, + 24, + ) + + x0 = ( + width + - nw + ) // 2 + y0 = ( + height + - nh + ) // 2 + + canvas[ + y0:y0 + nh, + x0:x0 + nw, + ] = resized + + return canvas + + +def make_status_panel( + probs: np.ndarray, + names: Dict[int, str], + width: int, + height: int, + pred_id: int, + gt_id: Optional[int], + gt_name: Optional[str], + label_conf: float, + infer_ms: float, + seg_metrics: Optional[dict], + group: str, + filename: str, +) -> np.ndarray: + panel = np.zeros( + ( + height, + width, + 3, + ), + dtype=np.uint8, + ) + panel[:] = ( + 29, + 29, + 29, + ) + + pred_name = names.get( + int(pred_id), + f"label_{pred_id}", + ) + + gt_display = ( + gt_name + or ( + names.get( + int(gt_id), + f"label_{gt_id}", + ) + if gt_id is not None + else "-" + ) + ) + + if gt_id is None: + verdict = "SEM GT" + else: + verdict = ( + "OK" + if int(gt_id) + == int(pred_id) + else "MISS" + ) + + header_lines = [ + f"{group}/{filename}", + ( + f"STATUS: {pred_name} ({pred_id}) " + f"conf={label_conf:.3f} | " + f"GT={gt_display} | {verdict}" + ), + f"infer={infer_ms:.2f} ms", + ] + + if seg_metrics is not None: + header_lines.append( + ( + f"SEG navIoU={seg_metrics['nav_iou']:.3f} " + f"nonNavIoU={seg_metrics['non_nav_iou']:.3f} " + f"unsafeNav={seg_metrics['unsafe_nav_rate']:.4f} " + f"acc={seg_metrics['acc']:.3f}" + ) + ) + + y = 28 + + for line in header_lines: + cv2.putText( + panel, + line, + ( + 14, + y, + ), + cv2.FONT_HERSHEY_SIMPLEX, + 0.58, + ( + 240, + 240, + 240, + ), + 1, + cv2.LINE_AA, + ) + y += 25 + + y += 8 + + # Ordena visualmente pelo ID para manter mapa mental estável. + n = len( + probs + ) + usable_h = max( + 1, + height + - y + - 10, + ) + row_h = max( + 24, + min( + 38, + usable_h + // max( + 1, + n, + ), + ), + ) + + label_w = min( + 260, + max( + 150, + width + // 5, + ), + ) + + bar_x0 = ( + label_w + + 24 + ) + bar_x1 = ( + width + - 90 + ) + bar_total = max( + 20, + bar_x1 + - bar_x0 + ) + + for i in range(n): + p = float( + probs[i] + ) + + name = names.get( + i, + f"label_{i}", + ) + + is_pred = ( + i + == int( + pred_id + ) + ) + is_gt = ( + gt_id is not None + and i + == int( + gt_id + ) + ) + + txt = f"{i} {name}" + + if is_pred: + txt += " < PRED" + if is_gt: + txt += " < GT" + + yy = ( + y + + i + * row_h + ) + + cv2.putText( + panel, + txt, + ( + 14, + yy + 18, + ), + cv2.FONT_HERSHEY_SIMPLEX, + 0.48, + ( + 245, + 245, + 245, + ), + 1, + cv2.LINE_AA, + ) + + cv2.rectangle( + panel, + ( + bar_x0, + yy + 5, + ), + ( + bar_x0 + + bar_total, + yy + row_h - 7, + ), + ( + 58, + 58, + 58, + ), + -1, + ) + + bw = int( + round( + bar_total + * clamp01( + p + ) + ) + ) + + # Sem codificar semanticamente por cor. Destaque vem do contorno/texto. + cv2.rectangle( + panel, + ( + bar_x0, + yy + 5, + ), + ( + bar_x0 + bw, + yy + row_h - 7, + ), + ( + 135, + 170, + 105, + ), + -1, + ) + + if is_pred: + cv2.rectangle( + panel, + ( + bar_x0 - 2, + yy + 3, + ), + ( + bar_x0 + + bar_total + + 2, + yy + row_h - 5, + ), + ( + 245, + 245, + 245, + ), + 1, + ) + + cv2.putText( + panel, + f"{p:.3f}", + ( + bar_x1 + 8, + yy + 18, + ), + cv2.FONT_HERSHEY_SIMPLEX, + 0.48, + ( + 235, + 235, + 235, + ), + 1, + cv2.LINE_AA, + ) + + return panel + + +def make_panel( + *, + sample: Sample, + img_bgr: np.ndarray, + gt_full: Optional[np.ndarray], + pred_full: np.ndarray, + classes: Dict[int, SegClass], + contract: RuntimeContract, + gt_label_id: Optional[int], + gt_label_name: Optional[str], + result: dict, + sample_seg_metrics: Optional[dict], + alpha: float, + max_width: int, +) -> np.ndarray: + pred_overlay = overlay_ids( + img_bgr, + pred_full, + classes, + alpha, + ) + + # Para pastas externas/legadas, tolera GT em resolução diferente da imagem. + # A métrica continua usando model-space; isto aqui é somente visualização. + gt_display = gt_full + + if ( + gt_display is not None + and gt_display.shape[:2] != img_bgr.shape[:2] + ): + gt_display = cv2.resize( + gt_display, + ( + img_bgr.shape[1], + img_bgr.shape[0], + ), + interpolation=cv2.INTER_NEAREST, + ) + + has_gt = ( + gt_display is not None + ) + + tiles: List[ + Tuple[str, np.ndarray] + ] = [ + ( + "ORIGINAL", + img_bgr, + ), + ] + + if has_gt: + gt_overlay = overlay_ids( + img_bgr, + gt_display, + classes, + alpha, + ) + + err = error_visual( + img_bgr, + gt_display, + pred_full, + contract.nav_class_id, + contract.non_nav_class_id, + ) + + tiles.extend([ + ( + "GT", + gt_overlay, + ), + ( + "PRED", + pred_overlay, + ), + ( + "ERRO | vermelho=falso navegável", + err, + ), + ]) + else: + pred_color = colorize_ids( + pred_full, + classes, + ) + + tiles.extend([ + ( + "PRED", + pred_overlay, + ), + ( + "MÁSCARA PRED", + pred_color, + ), + ]) + + cols = len( + tiles + ) + + tile_w = max( + 300, + int( + max_width + / cols + ), + ) + + aspect = ( + img_bgr.shape[0] + / max( + 1, + img_bgr.shape[1], + ) + ) + + tile_h = max( + 240, + int( + tile_w + * aspect + ), + ) + + top = np.concatenate( + [ + put_title( + resize_letterbox( + tile, + tile_w, + tile_h, + ), + title, + ) + for title, tile in tiles + ], + axis=1, + ) + + status_h = max( + 300, + 110 + + 34 + * len( + contract.label_name_by_id + ), + ) + + status = make_status_panel( + probs=result[ + "status_probs" + ], + names=contract.label_name_by_id, + width=top.shape[1], + height=status_h, + pred_id=result[ + "label_id" + ], + gt_id=gt_label_id, + gt_name=gt_label_name, + label_conf=result[ + "label_conf" + ], + infer_ms=result[ + "infer_ms" + ], + seg_metrics=sample_seg_metrics, + group=sample.group, + filename=sample.image.name, + ) + + return np.concatenate( + [ + top, + status, + ], + axis=0, + ) + + +# ============================================================================= +# Session metrics +# ============================================================================= + +class SessionMetrics: + def __init__( + self, + num_seg_classes: int, + num_label_classes: int, + nav_class_id: int, + non_nav_class_id: int, + ): + self.num_seg_classes = int( + num_seg_classes + ) + self.num_label_classes = int( + num_label_classes + ) + self.nav_class_id = int( + nav_class_id + ) + self.non_nav_class_id = int( + non_nav_class_id + ) + + self.records: Dict[ + str, + dict, + ] = {} + + def put( + self, + key: str, + record: dict, + ): + # Interactive pode voltar/avançar. Uma amostra conta uma única vez. + self.records[ + str(key) + ] = record + + def summary(self) -> dict: + cm_seg = np.zeros( + ( + self.num_seg_classes, + self.num_seg_classes, + ), + dtype=np.int64, + ) + cm_status = np.zeros( + ( + self.num_label_classes, + self.num_label_classes, + ), + dtype=np.int64, + ) + + infer_ms = [] + groups = Counter() + + for record in self.records.values(): + if ( + record.get( + "cm_seg" + ) + is not None + ): + cm_seg += record[ + "cm_seg" + ] + + gt_label_id = record.get( + "gt_label_id" + ) + pred_label_id = record.get( + "pred_label_id" + ) + + if ( + gt_label_id is not None + and pred_label_id is not None + and 0 + <= int( + gt_label_id + ) + < self.num_label_classes + and 0 + <= int( + pred_label_id + ) + < self.num_label_classes + ): + cm_status[ + int( + gt_label_id + ), + int( + pred_label_id + ), + ] += 1 + + infer_ms.append( + float( + record.get( + "infer_ms", + 0.0, + ) + ) + ) + groups[ + str( + record.get( + "group", + "unknown", + ) + ) + ] += 1 + + seg = ( + metrics_from_cm( + cm_seg, + self.nav_class_id, + self.non_nav_class_id, + ) + if cm_seg.sum() + > 0 + else None + ) + + status = ( + status_metrics_from_cm( + cm_status + ) + if cm_status.sum() + > 0 + else None + ) + + return { + "samples_unique": len( + self.records + ), + "groups": dict( + sorted( + groups.items() + ) + ), + "infer_ms_mean": ( + float( + np.mean( + infer_ms + ) + ) + if infer_ms + else None + ), + "infer_ms_p95": ( + float( + np.percentile( + infer_ms, + 95, + ) + ) + if infer_ms + else None + ), + "seg": seg, + "status": status, + "cm_seg": cm_seg.tolist(), + "cm_status": cm_status.tolist(), + } + + +def save_records_csv( + path: Path, + records: Dict[str, dict], +): + path.parent.mkdir( + parents=True, + exist_ok=True, + ) + + fields = [ + "base", + "group", + "image", + "infer_ms", + "gt_label_id", + "pred_label_id", + "label_conf", + "label_ok", + "seg_acc", + "nav_iou", + "non_nav_iou", + "unsafe_nav_rate", + ] + + with path.open( + "w", + newline="", + encoding="utf-8-sig", + ) as f: + writer = csv.DictWriter( + f, + fieldnames=fields, + ) + writer.writeheader() + + for record in records.values(): + writer.writerow({ + key: record.get( + key + ) + for key in fields + }) + + +# ============================================================================= +# Camera +# ============================================================================= + +class OakCamera: + """ + Compatibilidade simples com DepthAI moderno e legado. + + O objetivo é fornecer frame BGR 1080p ao mesmo preprocess usado pelos arquivos. + """ + + def __init__( + self, + fps: float, + resolution: str, + ): + self.fps = float( + fps + ) + self.resolution = str( + resolution + ).lower() + + self.mode = None + self.pipeline = None + self.device = None + self.queue = None + + def start(self): + try: + import depthai as dai + except ImportError as exc: + raise ImportError( + "Modo câmera exige depthai instalado." + ) from exc + + self.dai = dai + + target_size = ( + (1920, 1080) + if self.resolution == "1080p" + else ( + 1280, + 720, + ) + ) + + # API nova, igual ao estilo do script anterior. + try: + pipeline = dai.Pipeline() + + cam = pipeline.create( + dai.node.Camera + ).build() + + out = cam.requestOutput( + size=target_size, + type=dai.ImgFrame.Type.BGR888p, + fps=self.fps, + ) + + queue = out.createOutputQueue( + maxSize=4, + blocking=False, + ) + + pipeline.start() + + self.mode = "modern" + self.pipeline = pipeline + self.queue = queue + + print( + f"[CAM] DepthAI modern API | " + f"{target_size[0]}x{target_size[1]} @ {self.fps:g}" + ) + return self + + except Exception as modern_exc: + print( + f"[CAM] modern API indisponível: " + f"{type(modern_exc).__name__}: {modern_exc}" + ) + + # API clássica. + pipeline = dai.Pipeline() + + cam = pipeline.create( + dai.node.ColorCamera + ) + xout = pipeline.create( + dai.node.XLinkOut + ) + xout.setStreamName( + "rgb" + ) + + try: + cam.setBoardSocket( + dai.CameraBoardSocket.CAM_A + ) + except Exception: + pass + + if self.resolution == "1080p": + cam.setResolution( + dai.ColorCameraProperties.SensorResolution.THE_1080_P + ) + else: + cam.setResolution( + dai.ColorCameraProperties.SensorResolution.THE_720_P + ) + + cam.setFps( + self.fps + ) + cam.setInterleaved( + False + ) + + # getCvFrame entrega BGR ao host. + cam.video.link( + xout.input + ) + + device = dai.Device( + pipeline + ) + queue = device.getOutputQueue( + name="rgb", + maxSize=4, + blocking=False, + ) + + self.mode = "classic" + self.pipeline = pipeline + self.device = device + self.queue = queue + + print( + f"[CAM] DepthAI classic API | " + f"{target_size[0]}x{target_size[1]} @ {self.fps:g}" + ) + + return self + + def read(self) -> np.ndarray: + if self.queue is None: + raise RuntimeError( + "Câmera não inicializada." + ) + + if self.mode == "modern": + frame = self.queue.get() + else: + frame = self.queue.get() + + img = frame.getCvFrame() + + if img is None: + raise RuntimeError( + "Frame OAK vazio." + ) + + return img + + def running(self) -> bool: + if self.mode == "modern": + try: + return bool( + self.pipeline.isRunning() + ) + except Exception: + return True + + return True + + def close(self): + try: + if ( + self.mode == "modern" + and self.pipeline is not None + ): + self.pipeline.stop() + except Exception: + pass + + try: + if self.device is not None: + self.device.close() + except Exception: + pass + + +# ============================================================================= +# Console +# ============================================================================= + +def print_contract( + contract: RuntimeContract, + device: torch.device, +): + print("=" * 92) + print( + f"Agri Corridor Manual Test | {TESTER_VERSION}" + ) + print("=" * 92) + print( + f"Device : {device}" + ) + print( + f"Checkpoint : {contract.checkpoint_path}" + ) + print( + f"Epoch : {contract.checkpoint_epoch}" + ) + print( + f"Trainer : {contract.trainer_version}" + ) + print( + f"Backbone : {contract.backbone}" + ) + print( + f"Resolution : " + f"{contract.resolution_wh[0]}x{contract.resolution_wh[1]}" + ) + print( + f"ROI : " + f"begin={contract.roi_inicio:.3f} " + f"size={contract.roi_tamanho:.3f}" + ) + print( + f"Seg classes : {contract.seg_id2label}" + ) + print( + f"Nav / NonNav : " + f"{contract.nav_class_id} / {contract.non_nav_class_id}" + ) + print( + f"Status classes : {contract.label_name_by_id}" + ) + print( + f"Status head : {contract.status_head_cfg}" + ) + print( + f"Norm channels : {contract.norm_channels}" + ) + print( + f"Norm mean : {contract.norm_mean.tolist()}" + ) + print( + f"Norm std : {contract.norm_std.tolist()}" + ) + print("=" * 92) + + +def print_summary( + summary: dict, + contract: RuntimeContract, +): + print() + print("=" * 92) + print("RESUMO DA SESSÃO") + print("=" * 92) + print( + f"Amostras únicas : {summary['samples_unique']}" + ) + print( + f"Grupos : {summary['groups']}" + ) + + if ( + summary.get( + "infer_ms_mean" + ) + is not None + ): + print( + f"Inferência : " + f"mean={summary['infer_ms_mean']:.2f} ms " + f"p95={summary['infer_ms_p95']:.2f} ms" + ) + + seg = summary.get( + "seg" + ) + + if seg: + print( + f"SEG : " + f"acc={seg['acc']:.4f} " + f"mIoU={seg['miou']:.4f} " + f"navIoU={seg['nav_iou']:.4f} " + f"nonNavIoU={seg['non_nav_iou']:.4f} " + f"unsafeNav={seg['unsafe_nav_rate']:.5f} " + f"safety={seg['safety']:.5f}" + ) + + status = summary.get( + "status" + ) + + if status: + print( + f"STATUS : " + f"acc={status['acc']:.4f} " + f"macroF1={status['macro_f1']:.4f}" + ) + + for i, f1 in enumerate( + status[ + "f1_per_class" + ] + ): + print( + f" {i:2d} " + f"{contract.label_name_by_id.get(i, str(i)):28s} " + f"F1={float(f1):.4f} " + f"support={int(status['support_per_class'][i])}" + ) + + print("=" * 92) + + +# ============================================================================= +# Main +# ============================================================================= + +def main(): + ap = argparse.ArgumentParser( + description=( + "Teste manual OAK-D Lite: dataset, pasta de imagens ou câmera." + ) + ) + + ap.add_argument( + "--config", + default="config.json", + ) + + ap.add_argument( + "--ckpt", + default=None, + help="Checkpoint .pt explícito.", + ) + + ap.add_argument( + "--checkpoint", + default="operational", + choices=[ + "operational", + "nav", + "status", + "last", + ], + help=( + "Checkpoint padrão quando --ckpt não for informado." + ), + ) + + ap.add_argument( + "--split", + default="val", + choices=[ + "train", + "val", + "test", + ], + help=( + "Split padrão dentro de dataset/x/split." + ), + ) + + ap.add_argument( + "--input", + default=None, + help=( + "Pasta estruturada ou pasta simples de imagens. " + "Se omitido usa --split." + ), + ) + + ap.add_argument( + "--camera", + action="store_true", + help="Usa OAK-D Lite em tempo real.", + ) + + ap.add_argument( + "--camera-fps", + type=float, + default=20.0, + ) + + ap.add_argument( + "--camera-res", + default="1080p", + choices=[ + "1080p", + "720p", + ], + ) + + ap.add_argument( + "--alpha", + type=float, + default=DEFAULT_OVERLAY_ALPHA, + ) + + ap.add_argument( + "--max-width", + type=int, + default=1800, + help="Largura aproximada máxima do painel.", + ) + + ap.add_argument( + "--start", + type=int, + default=0, + ) + + ap.add_argument( + "--limit", + type=int, + default=0, + help="0 = todas.", + ) + + ap.add_argument( + "--no-view", + action="store_true", + help="Processa lote sem abrir janela.", + ) + + ap.add_argument( + "--save-dir", + default=None, + help="Salva painéis + resultados.", + ) + + args = ap.parse_args() + + if args.camera and args.input: + raise ValueError( + "Use --camera OU --input, não ambos." + ) + + config_path = Path( + args.config + ).resolve() + + if not config_path.is_file(): + raise FileNotFoundError( + config_path + ) + + config = safe_json_load( + config_path + ) + + project_root = find_project_root( + config_path + ) + + device = torch.device( + "cuda" + if torch.cuda.is_available() + else "cpu" + ) + + if args.ckpt: + checkpoint_path = Path( + args.ckpt + ).resolve() + else: + checkpoint_path = default_checkpoint_path( + project_root, + config, + args.checkpoint, + ) + + ( + base_model, + status_head, + contract, + checkpoint_raw, + ) = build_runtime_from_checkpoint( + checkpoint_path, + config, + device, + ) + + print_contract( + contract, + device, + ) + + labelmap_path = ( + project_root + / "dataset" + / "labelmap.txt" + ) + + classes = load_labelmap( + labelmap_path + ) + + # Sanity entre checkpoint e labelmap atual. + for cid, expected_name in contract.seg_id2label.items(): + got = classes.get( + int(cid) + ) + + if ( + got is not None + and got.name.lower() + != str( + expected_name + ).lower() + ): + print( + f"[WARN] labelmap atual id={cid} name={got.name!r} " + f"difere do checkpoint {expected_name!r}" + ) + + save_dir = ( + Path( + args.save_dir + ).resolve() + if args.save_dir + else None + ) + + if save_dir is not None: + save_dir.mkdir( + parents=True, + exist_ok=True, + ) + + session = SessionMetrics( + num_seg_classes=len( + contract.seg_id2label + ), + num_label_classes=len( + contract.label_name_by_id + ), + nav_class_id=contract.nav_class_id, + non_nav_class_id=contract.non_nav_class_id, + ) + + # ------------------------------------------------------------------------- + # Camera + # ------------------------------------------------------------------------- + + if args.camera: + if args.no_view: + raise ValueError( + "--camera com --no-view não faz sentido neste tester manual." + ) + + camera = OakCamera( + fps=args.camera_fps, + resolution=args.camera_res, + ).start() + + win = "Agrobot | Corridor Test | CAMERA" + cv2.namedWindow( + win, + cv2.WINDOW_NORMAL, + ) + + frame_index = 0 + fps_ema = None + last_t = time.perf_counter() + + try: + while camera.running(): + frame = camera.read() + + result = infer( + base_model, + status_head, + frame, + contract, + device, + ) + + now = time.perf_counter() + inst_fps = ( + 1.0 + / max( + 1e-6, + now - last_t, + ) + ) + last_t = now + + fps_ema = ( + inst_fps + if fps_ema is None + else ( + 0.85 + * fps_ema + + 0.15 + * inst_fps + ) + ) + + fake_sample = Sample( + base=f"camera_{frame_index:08d}", + image=Path( + f"camera_{frame_index:08d}.png" + ), + mask=None, + label=None, + group="camera", + source_root=project_root, + ) + + panel = make_panel( + sample=fake_sample, + img_bgr=frame, + gt_full=None, + pred_full=result[ + "pred_full" + ], + classes=classes, + contract=contract, + gt_label_id=None, + gt_label_name=None, + result=result, + sample_seg_metrics=None, + alpha=args.alpha, + max_width=args.max_width, + ) + + fps_text = ( + f"CAM FPS={fps_ema:.1f} | " + f"infer={result['infer_ms']:.2f} ms | " + f"Q/ESC sair | S salvar" + ) + + cv2.putText( + panel, + fps_text, + ( + 12, + panel.shape[0] - 12, + ), + cv2.FONT_HERSHEY_SIMPLEX, + 0.55, + ( + 245, + 245, + 245, + ), + 1, + cv2.LINE_AA, + ) + + cv2.imshow( + win, + panel, + ) + + key = ( + cv2.waitKey( + 1 + ) + & 0xFF + ) + + if key in ( + ord("q"), + ord("Q"), + 27, + ): + break + + if ( + key + in ( + ord("s"), + ord("S"), + ) + and save_dir + is not None + ): + out = ( + save_dir + / f"camera_{frame_index:08d}.jpg" + ) + cv2.imwrite( + str(out), + panel, + [ + cv2.IMWRITE_JPEG_QUALITY, + 92, + ], + ) + print( + f"[SAVE] {out}" + ) + + frame_index += 1 + + finally: + camera.close() + cv2.destroyAllWindows() + + return + + # ------------------------------------------------------------------------- + # Finite source + # ------------------------------------------------------------------------- + + W, H = contract.resolution_wh + resolution_root = ( + project_root + / "dataset" + / f"{W}x{H}" + ) + + if args.input: + input_root = Path( + args.input + ).resolve() + else: + input_root = ( + resolution_root + / "split" + / args.split + ) + + if ( + not input_root.is_dir() + and args.split + == "test" + ): + print( + "[WARN] split/test não existe; usando split/val." + ) + input_root = ( + resolution_root + / "split" + / "val" + ) + + samples, source_kind = discover_samples( + input_root + ) + + start = max( + 0, + min( + len( + samples + ) - 1, + int( + args.start + ), + ), + ) + + if args.limit > 0: + samples = samples[ + start: + start + + int( + args.limit + ) + ] + idx = 0 + else: + idx = start + + print( + f"[SOURCE] kind={source_kind} root={input_root}" + ) + print( + f"[SOURCE] samples={len(samples)}" + ) + + has_any_gt_mask = any( + s.mask is not None + for s in samples + ) + has_any_gt_label = any( + s.label is not None + for s in samples + ) + + print( + f"[SOURCE] GT mask={has_any_gt_mask} " + f"GT label={has_any_gt_label}" + ) + + win = "Agrobot | Corridor Test | FILES" + + if not args.no_view: + cv2.namedWindow( + win, + cv2.WINDOW_NORMAL, + ) + + while True: + if not samples: + break + + sample = samples[ + idx + ] + + img = ensure_bgr( + sample.image + ) + + gt_full = load_gt_mask( + sample.mask, + classes, + ) + + ( + gt_label_id, + gt_label_name, + ) = read_gt_label( + sample, + contract.label_name_by_id, + ) + + result = infer( + base_model, + status_head, + img, + contract, + device, + ) + + gt_model = gt_to_model_space( + gt_full, + result[ + "geometry" + ], + contract.resolution_wh, + ) + + cm_seg = None + sample_seg_metrics = None + + if gt_model is not None: + cm_seg = confusion_from_arrays( + result[ + "pred_model" + ], + gt_model, + num_classes=len( + contract.seg_id2label + ), + ) + + sample_seg_metrics = metrics_from_cm( + cm_seg, + contract.nav_class_id, + contract.non_nav_class_id, + ) + + pred_label_id = int( + result[ + "label_id" + ] + ) + + label_ok = ( + None + if gt_label_id is None + else ( + int( + gt_label_id + ) + == pred_label_id + ) + ) + + record = { + "base": sample.base, + "group": sample.group, + "image": str( + sample.image + ), + "infer_ms": float( + result[ + "infer_ms" + ] + ), + "gt_label_id": gt_label_id, + "pred_label_id": pred_label_id, + "label_conf": float( + result[ + "label_conf" + ] + ), + "label_ok": label_ok, + "cm_seg": cm_seg, + "seg_acc": ( + sample_seg_metrics[ + "acc" + ] + if sample_seg_metrics + else None + ), + "nav_iou": ( + sample_seg_metrics[ + "nav_iou" + ] + if sample_seg_metrics + else None + ), + "non_nav_iou": ( + sample_seg_metrics[ + "non_nav_iou" + ] + if sample_seg_metrics + else None + ), + "unsafe_nav_rate": ( + sample_seg_metrics[ + "unsafe_nav_rate" + ] + if sample_seg_metrics + else None + ), + } + + key_id = ( + f"{sample.group}/{sample.base}" + ) + + session.put( + key_id, + record, + ) + + panel = make_panel( + sample=sample, + img_bgr=img, + gt_full=gt_full, + pred_full=result[ + "pred_full" + ], + classes=classes, + contract=contract, + gt_label_id=gt_label_id, + gt_label_name=gt_label_name, + result=result, + sample_seg_metrics=sample_seg_metrics, + alpha=args.alpha, + max_width=args.max_width, + ) + + if save_dir is not None: + auto_save = bool( + args.no_view + ) + + if auto_save: + out = ( + save_dir + / ( + f"{idx:06d}_" + f"{sample.group.replace('/', '_')}_" + f"{sample.base}.jpg" + ) + ) + cv2.imwrite( + str(out), + panel, + [ + cv2.IMWRITE_JPEG_QUALITY, + 92, + ], + ) + + if args.no_view: + print( + f"[{idx + 1:5d}/{len(samples):5d}] " + f"{sample.group}/{sample.base} | " + f"status=" + f"{contract.label_name_by_id.get(pred_label_id, pred_label_id)} " + f"conf={result['label_conf']:.3f} " + f"infer={result['infer_ms']:.2f}ms" + + ( + ( + f" | navIoU={sample_seg_metrics['nav_iou']:.3f} " + f"unsafe={sample_seg_metrics['unsafe_nav_rate']:.4f}" + ) + if sample_seg_metrics + else "" + ) + + ( + ( + f" | label={'OK' if label_ok else 'MISS'}" + ) + if label_ok + is not None + else "" + ) + ) + + idx += 1 + + if idx >= len( + samples + ): + break + + continue + + cv2.imshow( + win, + panel, + ) + + key = ( + cv2.waitKey( + 0 + ) + & 0xFF + ) + + if key in ( + ord("q"), + ord("Q"), + 27, + ): + break + + if key in ( + ord("a"), + ord("A"), + 81, + ): + idx = ( + idx - 1 + ) % len( + samples + ) + continue + + if key in ( + ord("d"), + ord("D"), + 83, + ord(" "), + ): + idx = ( + idx + 1 + ) % len( + samples + ) + continue + + if key in ( + ord("s"), + ord("S"), + ): + target_dir = ( + save_dir + if save_dir + is not None + else ( + project_root + / "test_outputs" + ) + ) + target_dir.mkdir( + parents=True, + exist_ok=True, + ) + + out = ( + target_dir + / ( + f"{sample.group.replace('/', '_')}_" + f"{sample.base}.jpg" + ) + ) + cv2.imwrite( + str(out), + panel, + [ + cv2.IMWRITE_JPEG_QUALITY, + 92, + ], + ) + + print( + f"[SAVE] {out}" + ) + continue + + if not args.no_view: + cv2.destroyAllWindows() + + summary = session.summary() + + print_summary( + summary, + contract, + ) + + if save_dir is not None: + json_path = ( + save_dir + / "session_summary.json" + ) + + with json_path.open( + "w", + encoding="utf-8", + ) as f: + json.dump( + summary, + f, + ensure_ascii=False, + indent=2, + ) + + csv_path = ( + save_dir + / "results.csv" + ) + save_records_csv( + csv_path, + session.records, + ) + + print( + f"[SAVE] {json_path}" + ) + print( + f"[SAVE] {csv_path}" + ) + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/oak-d/_6_export_corridor_onnx.py b/Python/OAK/datasets/oak-d/_6_export_corridor_onnx.py new file mode 100644 index 000000000..af82ffd03 --- /dev/null +++ b/Python/OAK/datasets/oak-d/_6_export_corridor_onnx.py @@ -0,0 +1,3285 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +export_corridor_agri_onnx_v2.py +=============================== + +Exportador ONNX de CAMPO para o modelo frontal OAK-D Lite. + +Compatível com checkpoints gerados por: + train_corridor_agri_v2.py + +Contrato de campo v2 +-------------------- + +INPUT: + name : rgb_u8_nhwc + dtype : uint8 + layout : NHWC + shape : [1, H, W, 3] + color : RGB + range : [0, 255] + +O GRAFO FAZ: + uint8 -> float32 + NHWC -> NCHW + / 255 + normalize(mean/std do checkpoint) + SegFormer + CorridorStatusHead 3x4 + resize bilinear dos logits semânticos para HxW + argmax semântico + softmax do status + +OUTPUTS: + seg_ids + dtype : int32 + shape : [1, H, W] + + label_probs + dtype : float32 + shape : [1, K] + +Motivação: + O Visual Worker não precisa mais fazer em Python: + - astype(float32) + - /255 + - transpose HWC->CHW + - normalize + - resize dos logits + - argmax da segmentação + - softmax da cabeça global + + Ele só precisa: + 1) fornecer RGB uint8 na resolução do contrato; + 2) ler seg_ids; + 3) fazer argmax de apenas K probabilidades para obter o estado global. + +IMPORTANTE: + - batch fixo = 1; + - resolução fixa; + - checkpoint é a fonte da verdade; + - norm_stats externo NÃO é necessário no runtime; + - este exportador NÃO suporta o LabelHead antigo; + - este exportador NÃO exporta semantic_logits por padrão; + - AdaptiveAvgPool2d da StatusHead é convertido, SOMENTE no wrapper de + export fixo, para AvgPool2d estático matematicamente equivalente; + - a verificação PyTorch x ONNX usa a cabeça ORIGINAL do treinamento; + - o objetivo é o menor contrato útil e estável para campo. + +Exemplos: + + python export_corridor_agri_onnx_v2.py + + python export_corridor_agri_onnx_v2.py ^ + --checkpoint nav + + python export_corridor_agri_onnx_v2.py ^ + --ckpt backup/segformer_b0/nav_mit_corridor_agri_v2/best_operational.pt ^ + --out backup/segformer_b0/nav_mit_corridor_agri_v2/best_operational.field_v2.onnx + +Validação: + Após exportar, o script tenta: + - onnx.checker; + - shape inference; + - ONNX Runtime; + - comparação PyTorch x ONNX em input sintético, com tolerância OOD própria; + - comparação em algumas imagens reais de split/val, com tolerância mais rígida. + + Para testar especificamente com TensorRT EP: + --verify-provider tensorrt + + Para pular ONNX Runtime: + --no-ort-verify +""" + +from __future__ import annotations + +import argparse +import copy +import hashlib +import json +import os +import shutil +import time +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from typing import Dict, List, Optional, Sequence, Tuple + +import cv2 +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from transformers import SegformerForSemanticSegmentation + + +# ============================================================================= +# Constants +# ============================================================================= + +EXPORTER_VERSION = "agrobot_corridor_onnx_exporter_v2.1_fixed_pool" +CONTRACT_SCHEMA = "agrobot.visual.corridor.onnx.v2" +DEFAULT_OPSET = 17 + +INPUT_NAME = "rgb_u8_nhwc" +SEG_OUTPUT_NAME = "seg_ids" +STATUS_OUTPUT_NAME = "label_probs" + + +# ============================================================================= +# Generic helpers +# ============================================================================= + +def now_iso() -> str: + return datetime.now().isoformat( + timespec="seconds" + ) + + +def load_json( + path: Path, +) -> dict: + with path.open( + "r", + encoding="utf-8", + ) as f: + data = json.load( + f + ) + + if not isinstance( + data, + dict, + ): + raise ValueError( + f"Esperava objeto JSON em {path}" + ) + + return data + + +def save_json( + path: Path, + data: dict, +) -> None: + path.parent.mkdir( + parents=True, + exist_ok=True, + ) + + tmp = path.with_suffix( + path.suffix + + ".tmp" + ) + + with tmp.open( + "w", + encoding="utf-8", + ) as f: + json.dump( + data, + f, + ensure_ascii=False, + indent=2, + ) + + os.replace( + tmp, + path, + ) + + +def sha256_file( + path: Path, + chunk_size: int = 4 * 1024 * 1024, +) -> str: + h = hashlib.sha256() + + with path.open( + "rb" + ) as f: + while True: + block = f.read( + chunk_size + ) + + if not block: + break + + h.update( + block + ) + + return h.hexdigest() + + +def int_key_dict( + value, +) -> Dict[int, str]: + if not isinstance( + value, + dict, + ): + return {} + + return { + int(k): str(v) + for k, v in value.items() + } + + +def resolve_project_root( + config_path: Path, +) -> Path: + root = config_path.resolve().parent + + if not ( + root + / "dataset" + ).is_dir(): + raise FileNotFoundError( + f"dataset/ não encontrado ao lado de {config_path}. " + "Execute dentro da pasta oak-d ou passe --config correto." + ) + + return root + + +# ============================================================================= +# Status head, ESPELHO do trainer v2 +# ============================================================================= + +class CorridorStatusHead( + nn.Module +): + """ + Deve permanecer compatível com train_corridor_agri_v2.py. + + A cabeça: + - mantém feature profunda em resolução nativa; + - faz AdaptiveAvgPool espacial; + - resume separadamente feature e informação semântica; + - concatena vetores compactos; + - MLP final para estado global. + """ + + def __init__( + self, + feat_ch: int, + num_seg_classes: int, + num_label_classes: int, + pool_hw: Tuple[int, int], + hidden: int, + dropout: float, + detach_seg_summary: bool, + seg_summary: str, + ): + super().__init__() + + self.feat_ch = int( + feat_ch + ) + self.num_seg_classes = int( + num_seg_classes + ) + self.num_label_classes = int( + num_label_classes + ) + self.pool_hw = ( + int( + pool_hw[ + 0 + ] + ), + int( + pool_hw[ + 1 + ] + ), + ) + self.detach_seg_summary = bool( + detach_seg_summary + ) + self.seg_summary = str( + seg_summary + ).lower() + + pooled_cells = ( + self.pool_hw[ + 0 + ] + * self.pool_hw[ + 1 + ] + ) + + in_dim = ( + self.feat_ch + + self.num_seg_classes + ) * pooled_cells + + self.feat_pool = nn.AdaptiveAvgPool2d( + self.pool_hw + ) + self.seg_pool = nn.AdaptiveAvgPool2d( + self.pool_hw + ) + + self.net = nn.Sequential( + nn.Linear( + in_dim, + int( + hidden + ), + ), + nn.ReLU( + inplace=True + ), + nn.Dropout( + float( + dropout + ) + ), + nn.Linear( + int( + hidden + ), + self.num_label_classes, + ), + ) + + def forward( + self, + feat: torch.Tensor, + seg_logits_native: torch.Tensor, + ) -> torch.Tensor: + feat_grid = self.feat_pool( + feat + ).flatten( + 1 + ) + + seg_src = ( + seg_logits_native.detach() + if self.detach_seg_summary + else seg_logits_native + ) + + if ( + self.seg_summary + == "probabilities" + ): + seg_src = torch.softmax( + seg_src, + dim=1, + ) + elif ( + self.seg_summary + != "logits" + ): + raise RuntimeError( + f"seg_summary inválido: " + f"{self.seg_summary}" + ) + + seg_grid = self.seg_pool( + seg_src + ).flatten( + 1 + ) + + x = torch.cat( + [ + feat_grid, + seg_grid, + ], + dim=1, + ) + + return self.net( + x + ) + + +# ============================================================================= +# ONNX export adaptation for fixed-shape status pooling +# ============================================================================= + +@dataclass(frozen=True) +class StatusPoolExportPlan: + """ + Descreve a substituição export-only do AdaptiveAvgPool2d. + + O trainer continua usando AdaptiveAvgPool2d(pool_hw). + Para o ONNX de campo, cuja resolução é FIXA, quando a dimensão nativa é + exatamente divisível por pool_hw, o AdaptiveAvgPool2d é matematicamente + equivalente a AvgPool2d(kernel=stride=input/output). + + Isso evita a limitação do exporter TorchScript clássico, que pode perder + H/W estáticos depois de Softmax e falhar ao simbolizar adaptive_avg_pool2d. + """ + + feat_hw: Tuple[int, int] + seg_hw: Tuple[int, int] + pool_hw: Tuple[int, int] + + feat_kernel_hw: Tuple[int, int] + seg_kernel_hw: Tuple[int, int] + + method: str = "fixed_avgpool_exact" + + +def _exact_pool_kernel( + *, + input_hw: Tuple[int, int], + output_hw: Tuple[int, int], + name: str, +) -> Tuple[int, int]: + ih, iw = ( + int(input_hw[0]), + int(input_hw[1]), + ) + + oh, ow = ( + int(output_hw[0]), + int(output_hw[1]), + ) + + if ( + ih <= 0 + or iw <= 0 + or oh <= 0 + or ow <= 0 + ): + raise RuntimeError( + f"Dimensões inválidas para {name}: " + f"input={input_hw} output={output_hw}" + ) + + if ( + ih % oh != 0 + or iw % ow != 0 + ): + raise RuntimeError( + f"{name}: não é possível substituir AdaptiveAvgPool2d" + f"{output_hw} por AvgPool2d EXATO porque input={input_hw} " + "não é divisível célula-a-célula. " + "Não vou exportar uma aproximação silenciosa." + ) + + return ( + ih // oh, + iw // ow, + ) + + +def discover_status_native_shapes( + *, + base_model: nn.Module, + H: int, + W: int, + device: torch.device, + mean: Sequence[float], + std: Sequence[float], +) -> Tuple[ + Tuple[int, int], + Tuple[int, int], +]: + """ + Descobre as formas espaciais reais produzidas PELO checkpoint reconstruído. + + Não hardcodamos B0 nem 1024x576 aqui. Se backbone/resolução mudarem, o + exporter recalcula o plano e só aceita a adaptação quando ela for exata. + """ + dummy = torch.zeros( + ( + 1, + 3, + int(H), + int(W), + ), + dtype=torch.float32, + device=device, + ) + + mean_t = torch.tensor( + list(mean), + dtype=torch.float32, + device=device, + ).view( + 1, + 3, + 1, + 1, + ) + + std_t = torch.tensor( + list(std), + dtype=torch.float32, + device=device, + ).view( + 1, + 3, + 1, + 1, + ) + + with torch.inference_mode(): + out = base_model( + pixel_values=( + dummy + - mean_t + ) / std_t, + output_hidden_states=True, + return_dict=True, + ) + + seg_logits_native = out.logits + feat = out.hidden_states[ + -1 + ] + + feat_hw = ( + int(feat.shape[-2]), + int(feat.shape[-1]), + ) + + seg_hw = ( + int(seg_logits_native.shape[-2]), + int(seg_logits_native.shape[-1]), + ) + + return ( + feat_hw, + seg_hw, + ) + + +def build_export_status_head( + *, + reference_status_head: CorridorStatusHead, + feat_hw: Tuple[int, int], + seg_hw: Tuple[int, int], + device: torch.device, +) -> Tuple[ + CorridorStatusHead, + StatusPoolExportPlan, +]: + """ + Faz cópia SOMENTE da cabeça pequena. + + A cabeça original permanece intacta para a validação PyTorch de referência. + A cópia recebe AvgPool2d estático exato, amigável a ONNX/TensorRT. + """ + pool_hw = ( + int(reference_status_head.pool_hw[0]), + int(reference_status_head.pool_hw[1]), + ) + + feat_kernel = _exact_pool_kernel( + input_hw=feat_hw, + output_hw=pool_hw, + name="feat_pool", + ) + + seg_kernel = _exact_pool_kernel( + input_hw=seg_hw, + output_hw=pool_hw, + name="seg_pool", + ) + + export_head = copy.deepcopy( + reference_status_head + ) + + export_head.feat_pool = nn.AvgPool2d( + kernel_size=feat_kernel, + stride=feat_kernel, + padding=0, + ceil_mode=False, + count_include_pad=True, + ) + + export_head.seg_pool = nn.AvgPool2d( + kernel_size=seg_kernel, + stride=seg_kernel, + padding=0, + ceil_mode=False, + count_include_pad=True, + ) + + export_head.to( + device + ).eval() + + plan = StatusPoolExportPlan( + feat_hw=feat_hw, + seg_hw=seg_hw, + pool_hw=pool_hw, + feat_kernel_hw=feat_kernel, + seg_kernel_hw=seg_kernel, + ) + + return ( + export_head, + plan, + ) + + +def verify_export_wrapper_equivalence( + *, + reference_wrapper: nn.Module, + export_wrapper: nn.Module, + H: int, + W: int, + device: torch.device, + max_prob_abs: float = 1e-6, +) -> dict: + """ + Confirma ANTES do ONNX que a adaptação export-only preserva a função. + + A segmentação precisa ser bit-a-bit idêntica. + As probabilidades de status podem ter erro FP minúsculo por ordem de soma. + """ + gen = torch.Generator( + device=device + ) + + gen.manual_seed( + 20260915 + ) + + dummy = torch.randint( + low=0, + high=256, + size=( + 1, + int(H), + int(W), + 3, + ), + dtype=torch.uint8, + device=device, + generator=gen, + ) + + reference_wrapper.eval() + export_wrapper.eval() + + with torch.inference_mode(): + ref_seg, ref_probs = reference_wrapper( + dummy + ) + + exp_seg, exp_probs = export_wrapper( + dummy + ) + + seg_equal = bool( + torch.equal( + ref_seg, + exp_seg, + ) + ) + + prob_abs = torch.abs( + ref_probs.float() + - exp_probs.float() + ) + + prob_max_abs = float( + prob_abs.max().item() + ) + + prob_mean_abs = float( + prob_abs.mean().item() + ) + + label_ref = int( + torch.argmax( + ref_probs[ + 0 + ] + ).item() + ) + + label_exp = int( + torch.argmax( + exp_probs[ + 0 + ] + ).item() + ) + + passed = ( + seg_equal + and label_ref + == label_exp + and prob_max_abs + <= float( + max_prob_abs + ) + ) + + result = { + "passed": bool( + passed + ), + "seg_equal": bool( + seg_equal + ), + "label_same": bool( + label_ref + == label_exp + ), + "label_reference": label_ref, + "label_export": label_exp, + "status_prob_max_abs": prob_max_abs, + "status_prob_mean_abs": prob_mean_abs, + "max_prob_abs_allowed": float( + max_prob_abs + ), + } + + print( + "[EXPORT-ADAPT] PT original x PT export-friendly | " + f"seg_equal={seg_equal} " + f"label={label_ref}/{label_exp} " + f"pMaxAbs={prob_max_abs:.9g} " + f"{'OK' if passed else 'FAIL'}" + ) + + if not passed: + raise RuntimeError( + "A adaptação export-only do status head não preservou " + f"a saída dentro da tolerância: {result}" + ) + + return result + + +# ============================================================================= +# Checkpoint contract +# ============================================================================= + +@dataclass +class CheckpointContract: + checkpoint_path: Path + trainer_version: str + epoch: int + bests: dict + + backbone: str + resolution_wh: Tuple[int, int] + + roi_inicio: float + roi_tamanho: float + + seg_id2label: Dict[int, str] + nav_class_id: int + non_nav_class_id: int + + label_name_by_id: Dict[int, str] + group_name_by_id: Dict[int, str] + + norm_channels: List[str] + norm_mean: List[float] + norm_std: List[float] + + status_head: dict + checkpoint_metadata: dict + + +def default_checkpoint_path( + project_root: Path, + config: dict, + which: str, +) -> Path: + model_family = str( + config.get( + "modelo", + "segformer_b0", + ) + ) + + model_name = str( + config.get( + "model_name", + "nav_mit", + ) + ) + + save_dir = ( + project_root + / "backup" + / model_family + / f"{model_name}_corridor_agri_v2" + ) + + filename = { + "operational": "best_operational.pt", + "nav": "best_nav.pt", + "status": "best_status.pt", + "last": "last.pt", + }[ + which + ] + + return ( + save_dir + / filename + ) + + +def load_checkpoint_contract( + checkpoint_path: Path, + config: dict, +) -> Tuple[dict, CheckpointContract]: + if not checkpoint_path.is_file(): + raise FileNotFoundError( + f"Checkpoint não encontrado: {checkpoint_path}" + ) + + ckpt = torch.load( + checkpoint_path, + map_location="cpu", + weights_only=False, + ) + + if "model" not in ckpt: + raise RuntimeError( + "Checkpoint sem chave 'model'." + ) + + if "status_head" not in ckpt: + raise RuntimeError( + "Checkpoint sem chave 'status_head'. " + "O export v2 não aceita o antigo aux_head/LabelHead." + ) + + metadata = ckpt.get( + "metadata", + {}, + ) or {} + + trainer_version = str( + ckpt.get( + "trainer_version", + metadata.get( + "trainer_version", + "unknown", + ), + ) + ) + + resolution = metadata.get( + "resolution_wh", + config.get( + "resolucao", + [1024, 576], + ), + ) + + W = int( + resolution[ + 0 + ] + ) + H = int( + resolution[ + 1 + ] + ) + + backbone = str( + metadata.get( + "backbone", + config.get( + "backbone", + "nvidia/mit-b0", + ), + ) + ) + + seg_id2label = int_key_dict( + metadata.get( + "seg_id2label", + {}, + ) + ) + + if not seg_id2label: + raise RuntimeError( + "Checkpoint v2 sem metadata.seg_id2label." + ) + + nav_class_id = int( + metadata.get( + "nav_class_id" + ) + ) + non_nav_class_id = int( + metadata.get( + "non_nav_class_id" + ) + ) + + label_name_by_id = int_key_dict( + metadata.get( + "label_name_by_id", + {}, + ) + ) + + if not label_name_by_id: + raise RuntimeError( + "Checkpoint v2 sem metadata.label_name_by_id." + ) + + group_name_by_id = int_key_dict( + metadata.get( + "group_name_by_id", + {}, + ) + ) + + norm_channels = [ + str( + x + ).upper() + for x in metadata.get( + "norm_channels", + [], + ) + ] + + norm_mean = [ + float( + x + ) + for x in metadata.get( + "norm_mean", + [], + ) + ] + + norm_std = [ + float( + x + ) + for x in metadata.get( + "norm_std", + [], + ) + ] + + if ( + norm_channels + != [ + "R", + "G", + "B", + ] + ): + raise RuntimeError( + f"Checkpoint precisa norm_channels RGB; " + f"recebido={norm_channels}" + ) + + if ( + len( + norm_mean + ) + != 3 + or len( + norm_std + ) + != 3 + ): + raise RuntimeError( + f"Checkpoint possui norm stats inválidos: " + f"mean={norm_mean} std={norm_std}" + ) + + if ( + not np.all( + np.isfinite( + np.asarray( + norm_mean + + norm_std, + dtype=np.float64, + ) + ) + ) + or any( + s <= 0 + for s in norm_std + ) + ): + raise RuntimeError( + "Checkpoint possui mean/std inválidos." + ) + + status_meta = dict( + metadata.get( + "status_head", + {}, + ) + or {} + ) + + if not status_meta: + raise RuntimeError( + "Checkpoint v2 sem metadata.status_head." + ) + + target_contract = ( + metadata.get( + "onnx_target_contract", + {}, + ) + or {} + ) + + if target_contract: + expected_input = str( + target_contract.get( + "input", + "rgb_uint8_nhwc", + ) + ) + + if ( + expected_input + != "rgb_uint8_nhwc" + ): + raise RuntimeError( + f"Checkpoint pede ONNX input " + f"{expected_input!r}, mas exporter v2 " + "implementa rgb_uint8_nhwc." + ) + + contract = CheckpointContract( + checkpoint_path=checkpoint_path, + trainer_version=trainer_version, + epoch=int( + ckpt.get( + "epoch", + -1, + ) + ), + bests=dict( + ckpt.get( + "bests", + {}, + ) + or {} + ), + backbone=backbone, + resolution_wh=( + W, + H, + ), + roi_inicio=float( + metadata.get( + "roi_inicio", + config.get( + "roi_inicio", + 0.0, + ), + ) + ), + roi_tamanho=float( + metadata.get( + "roi_tamanho", + config.get( + "roi_tamanho", + 1.0, + ), + ) + ), + seg_id2label=seg_id2label, + nav_class_id=nav_class_id, + non_nav_class_id=non_nav_class_id, + label_name_by_id=label_name_by_id, + group_name_by_id=group_name_by_id, + norm_channels=norm_channels, + norm_mean=norm_mean, + norm_std=norm_std, + status_head=status_meta, + checkpoint_metadata=metadata, + ) + + return ( + ckpt, + contract, + ) + + +# ============================================================================= +# Build PyTorch runtime +# ============================================================================= + +def infer_hidden_from_status_state( + status_state: dict, +) -> int: + key = "net.0.weight" + + if key not in status_state: + raise RuntimeError( + f"status_head state_dict sem {key}" + ) + + return int( + status_state[ + key + ].shape[ + 0 + ] + ) + + +def infer_num_status_classes( + status_state: dict, +) -> int: + for key in ( + "net.3.weight", + "net.4.weight", + ): + if key in status_state: + return int( + status_state[ + key + ].shape[ + 0 + ] + ) + + raise RuntimeError( + "Não consegui inferir saída do status_head." + ) + + +def discover_feat_ch( + base_model: nn.Module, + H: int, + W: int, + device: torch.device, + mean: Sequence[float], + std: Sequence[float], +) -> int: + dummy = torch.zeros( + ( + 1, + 3, + H, + W, + ), + dtype=torch.float32, + device=device, + ) + + mean_t = torch.tensor( + mean, + dtype=torch.float32, + device=device, + ).view( + 1, + 3, + 1, + 1, + ) + + std_t = torch.tensor( + std, + dtype=torch.float32, + device=device, + ).view( + 1, + 3, + 1, + 1, + ) + + with torch.inference_mode(): + out = base_model( + pixel_values=( + dummy + - mean_t + ) / std_t, + output_hidden_states=True, + return_dict=True, + ) + + return int( + out.hidden_states[ + -1 + ].shape[ + 1 + ] + ) + + +def build_models( + ckpt: dict, + contract: CheckpointContract, + device: torch.device, +) -> Tuple[ + SegformerForSemanticSegmentation, + CorridorStatusHead, +]: + num_seg_classes = len( + contract.seg_id2label + ) + + num_status_classes = len( + contract.label_name_by_id + ) + + status_num_from_state = infer_num_status_classes( + ckpt[ + "status_head" + ] + ) + + if ( + status_num_from_state + != num_status_classes + ): + raise RuntimeError( + f"status classes metadata={num_status_classes}, " + f"state_dict={status_num_from_state}" + ) + + base_model = ( + SegformerForSemanticSegmentation + .from_pretrained( + contract.backbone, + num_labels=num_seg_classes, + ignore_mismatched_sizes=True, + use_safetensors=True, + ) + ) + + base_model.config.output_hidden_states = True + + base_model.load_state_dict( + ckpt[ + "model" + ], + strict=True, + ) + + base_model.to( + device + ).eval() + + status_meta = contract.status_head + + feat_ch = int( + status_meta.get( + "feat_ch", + 0, + ) + ) + + W, H = contract.resolution_wh + + if feat_ch <= 0: + feat_ch = discover_feat_ch( + base_model, + H, + W, + device, + contract.norm_mean, + contract.norm_std, + ) + + pool_hw_raw = status_meta.get( + "pool_hw" + ) + + if pool_hw_raw is None: + # Trainer salva pool_h/pool_w no config e pool_hw no export_config. + pool_h = int( + status_meta.get( + "pool_h", + 3, + ) + ) + pool_w = int( + status_meta.get( + "pool_w", + 4, + ) + ) + pool_hw = ( + pool_h, + pool_w, + ) + else: + pool_hw = ( + int( + pool_hw_raw[ + 0 + ] + ), + int( + pool_hw_raw[ + 1 + ] + ), + ) + + hidden = infer_hidden_from_status_state( + ckpt[ + "status_head" + ] + ) + + status_head = CorridorStatusHead( + feat_ch=feat_ch, + num_seg_classes=num_seg_classes, + num_label_classes=num_status_classes, + pool_hw=pool_hw, + hidden=hidden, + dropout=0.0, + detach_seg_summary=bool( + status_meta.get( + "detach_seg_summary", + True, + ) + ), + seg_summary=str( + status_meta.get( + "seg_summary", + "probabilities", + ) + ), + ) + + status_head.load_state_dict( + ckpt[ + "status_head" + ], + strict=True, + ) + + status_head.to( + device + ).eval() + + return ( + base_model, + status_head, + ) + + +# ============================================================================= +# Field ONNX wrapper +# ============================================================================= + +class CorridorFieldOnnx( + nn.Module +): + """ + CONTRATO FIXO DE CAMPO. + + Entrada: + RGB uint8 NHWC na resolução do modelo. + + Saída: + seg_ids int32 HxW + label_probs float32 K + """ + + def __init__( + self, + base_model: nn.Module, + status_head: CorridorStatusHead, + resolution_wh: Tuple[int, int], + norm_mean: Sequence[float], + norm_std: Sequence[float], + ): + super().__init__() + + self.base_model = base_model + self.status_head = status_head + + self.W = int( + resolution_wh[ + 0 + ] + ) + self.H = int( + resolution_wh[ + 1 + ] + ) + + mean = torch.tensor( + list( + norm_mean + ), + dtype=torch.float32, + ).view( + 1, + 3, + 1, + 1, + ) + + std = torch.tensor( + list( + norm_std + ), + dtype=torch.float32, + ).view( + 1, + 3, + 1, + 1, + ) + + self.register_buffer( + "norm_mean", + mean, + ) + self.register_buffer( + "norm_std", + std, + ) + + def forward( + self, + rgb_u8_nhwc: torch.Tensor, + ): + # uint8 network I/O -> float. + x = rgb_u8_nhwc.to( + dtype=torch.float32 + ) + + # [0,255] -> [0,1] + x = x / 255.0 + + # NHWC -> NCHW + x = x.permute( + 0, + 3, + 1, + 2, + ) + + # Normalização EXATA do treinamento. + x = ( + x + - self.norm_mean + ) / self.norm_std + + out = self.base_model( + pixel_values=x, + output_hidden_states=True, + return_dict=True, + ) + + seg_logits_native = out.logits + feat = out.hidden_states[ + -1 + ] + + status_logits = self.status_head( + feat, + seg_logits_native, + ) + + # Treinamento/validação fazem resize dos LOGITS e só depois argmax. + seg_logits_full = F.interpolate( + seg_logits_native, + size=( + self.H, + self.W, + ), + mode="bilinear", + align_corners=False, + ) + + # ArgMax ONNX nasce como int64. Cast explícito para int32 deixa + # o output menor e adequado ao runtime TensorRT. + seg_ids = torch.argmax( + seg_logits_full, + dim=1, + ).to( + torch.int32 + ) + + label_probs = torch.softmax( + status_logits, + dim=1, + ) + + return ( + seg_ids, + label_probs, + ) + + +# ============================================================================= +# Input preparation for verification ONLY +# ============================================================================= + +def image_to_contract_input( + image_path: Path, + contract: CheckpointContract, +) -> np.ndarray: + """ + Converte um arquivo qualquer para o tensor EXTERNO esperado pelo ONNX. + + Isto é só para verificação do export. + + O ONNX NÃO contém crop/resize geométrico de câmera: + - aplica ROI da source; + - resize para W,H; + - converte BGR OpenCV -> RGB; + - mantém uint8 NHWC. + """ + bgr = cv2.imread( + str( + image_path + ), + cv2.IMREAD_COLOR, + ) + + if bgr is None: + raise RuntimeError( + f"Falha ao ler imagem: {image_path}" + ) + + h, w = bgr.shape[ + :2 + ] + + y0 = int( + round( + h + * contract.roi_inicio + ) + ) + + y1 = int( + round( + h + * min( + 1.0, + contract.roi_inicio + + contract.roi_tamanho, + ) + ) + ) + + y0 = max( + 0, + min( + h - 1, + y0, + ), + ) + + y1 = max( + y0 + 1, + min( + h, + y1, + ), + ) + + roi = bgr[ + y0:y1, + :, + ] + + W, H = contract.resolution_wh + + resized = cv2.resize( + roi, + ( + W, + H, + ), + interpolation=cv2.INTER_AREA, + ) + + rgb = cv2.cvtColor( + resized, + cv2.COLOR_BGR2RGB, + ) + + return np.ascontiguousarray( + rgb[ + None, + ..., + ], + dtype=np.uint8, + ) + + +def find_verification_images( + project_root: Path, + contract: CheckpointContract, + explicit_dir: Optional[Path], + count: int, +) -> List[Path]: + if count <= 0: + return [] + + if explicit_dir is not None: + roots = [ + explicit_dir + ] + else: + W, H = contract.resolution_wh + + roots = [ + ( + project_root + / "dataset" + / f"{W}x{H}" + / "split" + / "val" + / "images" + ), + ( + project_root + / "dataset" + / f"{W}x{H}" + / "group" + ), + ] + + candidates: List[ + Path + ] = [] + + exts = { + ".png", + ".jpg", + ".jpeg", + ".bmp", + ".webp", + } + + for root in roots: + if not root.is_dir(): + continue + + if ( + root.name + == "images" + ): + found = [ + p + for p in root.iterdir() + if ( + p.is_file() + and p.suffix.lower() + in exts + ) + ] + else: + found = [ + p + for p in root.rglob( + "*" + ) + if ( + p.is_file() + and p.parent.name + == "images" + and p.suffix.lower() + in exts + ) + ] + + candidates.extend( + found + ) + + if candidates: + break + + candidates = sorted( + set( + candidates + ), + key=lambda p: str( + p + ).lower(), + ) + + if not candidates: + return [] + + if len( + candidates + ) <= count: + return candidates + + idxs = np.linspace( + 0, + len( + candidates + ) - 1, + count, + ).astype( + int + ) + + return [ + candidates[ + int( + i + ) + ] + for i in idxs + ] + + +# ============================================================================= +# ONNX metadata +# ============================================================================= + +def contract_core_dict( + *, + checkpoint_contract: CheckpointContract, + checkpoint_sha256: str, + opset: int, + export_pool_plan: Optional[ + StatusPoolExportPlan + ] = None, +) -> dict: + W, H = checkpoint_contract.resolution_wh + + return { + "schema": CONTRACT_SCHEMA, + "contract_version": 2, + "exporter_version": EXPORTER_VERSION, + "created_at": now_iso(), + + "model": { + "kind": "agrobot_visual_corridor", + "task": "semantic_navigation_plus_corridor_status", + "backbone": checkpoint_contract.backbone, + "trainer_version": checkpoint_contract.trainer_version, + "checkpoint_epoch": checkpoint_contract.epoch, + "checkpoint_bests": checkpoint_contract.bests, + "checkpoint_sha256": checkpoint_sha256, + }, + + "input": { + "name": INPUT_NAME, + "dtype": "uint8", + "shape": [ + 1, + H, + W, + 3, + ], + "layout": "NHWC", + "color_order": "RGB", + "value_range": [ + 0, + 255, + ], + "batch": 1, + "geometry": { + "input_is_already_model_resolution": True, + "source_roi_inicio": checkpoint_contract.roi_inicio, + "source_roi_tamanho": checkpoint_contract.roi_tamanho, + "source_to_model_resize": "INTER_AREA_recommended", + "model_resolution_wh": [ + W, + H, + ], + }, + }, + + "preprocess_in_graph": { + "enabled": True, + "cast": "uint8_to_float32", + "scale": "x/255.0", + "layout": "NHWC_to_NCHW", + "normalize": "(x-mean)/std", + "channels": checkpoint_contract.norm_channels, + "mean": checkpoint_contract.norm_mean, + "std": checkpoint_contract.norm_std, + }, + + "outputs": { + SEG_OUTPUT_NAME: { + "name": SEG_OUTPUT_NAME, + "dtype": "int32", + "shape": [ + 1, + H, + W, + ], + "meaning": "semantic_class_id", + "postprocess_in_graph": ( + "native_logits -> bilinear_resize_to_input -> argmax -> int32" + ), + "id2label": { + str( + k + ): v + for k, v in checkpoint_contract.seg_id2label.items() + }, + "nav_class_id": checkpoint_contract.nav_class_id, + "non_nav_class_id": checkpoint_contract.non_nav_class_id, + }, + + STATUS_OUTPUT_NAME: { + "name": STATUS_OUTPUT_NAME, + "dtype": "float32", + "shape": [ + 1, + len( + checkpoint_contract.label_name_by_id + ), + ], + "meaning": "softmax_probability_per_corridor_state", + "decision_rule": "argmax(label_probs[0])", + "id2label": { + str( + k + ): v + for k, v in checkpoint_contract.label_name_by_id.items() + }, + }, + }, + + "status_head": checkpoint_contract.status_head, + + "export_adaptation": ( + { + "status_pool": { + "method": export_pool_plan.method, + "training_operator": "AdaptiveAvgPool2d", + "onnx_operator": "AveragePool", + "mathematically_exact_for_fixed_shape": True, + "feat_hw": list( + export_pool_plan.feat_hw + ), + "seg_hw": list( + export_pool_plan.seg_hw + ), + "pool_hw": list( + export_pool_plan.pool_hw + ), + "feat_kernel_hw": list( + export_pool_plan.feat_kernel_hw + ), + "seg_kernel_hw": list( + export_pool_plan.seg_kernel_hw + ), + } + } + if export_pool_plan + is not None + else {} + ), + + "runtime_contract": { + "fixed_batch": True, + "fixed_spatial_shape": True, + "dynamic_axes": False, + "external_normalization_required": False, + "external_semantic_argmax_required": False, + "external_semantic_resize_required": False, + "external_status_softmax_required": False, + "only_external_status_argmax_required": True, + }, + + "onnx": { + "opset": int( + opset + ), + "input_names": [ + INPUT_NAME + ], + "output_names": [ + SEG_OUTPUT_NAME, + STATUS_OUTPUT_NAME, + ], + }, + } + + +def embed_onnx_metadata( + onnx_path: Path, + core_contract: dict, +) -> bool: + try: + import onnx + except ImportError: + return False + + model = onnx.load( + str( + onnx_path + ) + ) + + del model.metadata_props[ + : + ] + + metadata = { + "agrobot.schema": CONTRACT_SCHEMA, + "agrobot.contract_version": "2", + "agrobot.exporter_version": EXPORTER_VERSION, + "agrobot.input_name": INPUT_NAME, + "agrobot.input_dtype": "uint8", + "agrobot.input_layout": "NHWC", + "agrobot.input_color_order": "RGB", + "agrobot.seg_output_name": SEG_OUTPUT_NAME, + "agrobot.status_output_name": STATUS_OUTPUT_NAME, + "agrobot.contract_json": json.dumps( + core_contract, + ensure_ascii=False, + separators=( + ",", + ":", + ), + ), + } + + for key, value in metadata.items(): + prop = model.metadata_props.add() + prop.key = str( + key + ) + prop.value = str( + value + ) + + onnx.save( + model, + str( + onnx_path + ), + ) + + return True + + +# ============================================================================= +# Export +# ============================================================================= + +def export_onnx( + wrapper: CorridorFieldOnnx, + out_path: Path, + H: int, + W: int, + device: torch.device, + opset: int, +) -> None: + wrapper.to( + device + ).eval() + + # Input real do campo é uint8. + dummy = torch.randint( + low=0, + high=256, + size=( + 1, + H, + W, + 3, + ), + dtype=torch.uint8, + device=device, + ) + + with torch.inference_mode(): + seg_ids, label_probs = wrapper( + dummy + ) + + print( + f"[CHECK] {INPUT_NAME}: " + f"shape={tuple(dummy.shape)} dtype={dummy.dtype}" + ) + print( + f"[CHECK] {SEG_OUTPUT_NAME}: " + f"shape={tuple(seg_ids.shape)} dtype={seg_ids.dtype}" + ) + print( + f"[CHECK] {STATUS_OUTPUT_NAME}: " + f"shape={tuple(label_probs.shape)} dtype={label_probs.dtype}" + ) + + out_path.parent.mkdir( + parents=True, + exist_ok=True, + ) + + kwargs = dict( + model=wrapper, + args=( + dummy, + ), + f=str( + out_path + ), + export_params=True, + opset_version=int( + opset + ), + do_constant_folding=True, + input_names=[ + INPUT_NAME + ], + output_names=[ + SEG_OUTPUT_NAME, + STATUS_OUTPUT_NAME, + ], + dynamic_axes=None, + ) + + print( + "[EXPORT] torch.onnx.export " + f"opset={opset} fixed_batch=1 fixed_shape={H}x{W}" + ) + + # Para TensorRT 10.x, o exporter clássico continua uma ponte previsível. + # Em PyTorch novo, dynamo pode ser default; forçamos legacy quando aceito. + try: + torch.onnx.export( + **kwargs, + dynamo=False, + ) + except TypeError: + torch.onnx.export( + **kwargs + ) + + +# ============================================================================= +# ONNX checks +# ============================================================================= + +def check_onnx( + path: Path, +) -> dict: + result = { + "onnx_installed": False, + "checker_ok": None, + "shape_inference_ok": None, + "inputs": [], + "outputs": [], + "error": None, + } + + try: + import onnx + except ImportError: + print( + "[WARN] pacote onnx não instalado; " + "checker/shape inference pulados." + ) + return result + + result[ + "onnx_installed" + ] = True + + try: + model = onnx.load( + str( + path + ) + ) + + onnx.checker.check_model( + model + ) + + result[ + "checker_ok" + ] = True + + try: + inferred = ( + onnx.shape_inference + .infer_shapes( + model + ) + ) + + result[ + "shape_inference_ok" + ] = True + + graph = inferred.graph + + except Exception as exc: + result[ + "shape_inference_ok" + ] = False + graph = model.graph + + print( + f"[WARN] ONNX shape inference: {exc}" + ) + + result[ + "inputs" + ] = [ + x.name + for x in graph.input + ] + + result[ + "outputs" + ] = [ + x.name + for x in graph.output + ] + + print( + "[OK] onnx.checker passou." + ) + print( + f"[ONNX] inputs={result['inputs']}" + ) + print( + f"[ONNX] outputs={result['outputs']}" + ) + + except Exception as exc: + result[ + "checker_ok" + ] = False + result[ + "error" + ] = ( + f"{type(exc).__name__}: {exc}" + ) + + raise + + return result + + +def create_ort_session( + onnx_path: Path, + provider_mode: str, +): + try: + import onnxruntime as ort + except ImportError: + return None, { + "installed": False, + "available": [], + "requested": provider_mode, + "active": [], + } + + available = ort.get_available_providers() + + provider_mode = str( + provider_mode + ).lower() + + if provider_mode == "auto": + if ( + "CUDAExecutionProvider" + in available + ): + providers = [ + "CUDAExecutionProvider", + "CPUExecutionProvider", + ] + else: + providers = [ + "CPUExecutionProvider" + ] + + elif provider_mode == "cuda": + providers = [ + "CUDAExecutionProvider", + "CPUExecutionProvider", + ] + + elif provider_mode == "cpu": + providers = [ + "CPUExecutionProvider" + ] + + elif provider_mode == "tensorrt": + cache_dir = ( + onnx_path.parent + / "trt_cache_field_v2" + ) + + cache_dir.mkdir( + parents=True, + exist_ok=True, + ) + + trt_options = { + "trt_engine_cache_enable": True, + "trt_engine_cache_path": str( + cache_dir + ), + "trt_timing_cache_enable": True, + "trt_timing_cache_path": str( + cache_dir + ), + "trt_fp16_enable": True, + } + + providers = [ + ( + "TensorrtExecutionProvider", + trt_options, + ), + "CUDAExecutionProvider", + "CPUExecutionProvider", + ] + + else: + raise ValueError( + provider_mode + ) + + providers_ok = [] + + for p in providers: + name = ( + p[ + 0 + ] + if isinstance( + p, + tuple, + ) + else p + ) + + if name in available: + providers_ok.append( + p + ) + + if not providers_ok: + raise RuntimeError( + f"Nenhum provider ORT utilizável. " + f"requested={provider_mode} available={available}" + ) + + options = ort.SessionOptions() + options.graph_optimization_level = ( + ort.GraphOptimizationLevel + .ORT_ENABLE_ALL + ) + + session = ort.InferenceSession( + str( + onnx_path + ), + sess_options=options, + providers=providers_ok, + ) + + active = session.get_providers() + + if ( + provider_mode + == "cuda" + and "CUDAExecutionProvider" + not in active + ): + raise RuntimeError( + f"CUDA solicitado, mas providers ativos={active}" + ) + + if ( + provider_mode + == "tensorrt" + and "TensorrtExecutionProvider" + not in active + ): + raise RuntimeError( + f"TensorRT solicitado, mas providers ativos={active}" + ) + + info = { + "installed": True, + "available": list( + available + ), + "requested": provider_mode, + "active": list( + active + ), + } + + return ( + session, + info, + ) + + +def run_pytorch_wrapper( + wrapper: CorridorFieldOnnx, + input_np: np.ndarray, + device: torch.device, +): + x = torch.from_numpy( + np.ascontiguousarray( + input_np + ) + ).to( + device + ) + + with torch.inference_mode(): + seg, probs = wrapper( + x + ) + + return ( + seg.detach() + .cpu() + .numpy(), + probs.detach() + .cpu() + .numpy(), + ) + + +def compare_outputs( + pt_seg: np.ndarray, + pt_probs: np.ndarray, + ort_seg: np.ndarray, + ort_probs: np.ndarray, +) -> dict: + if ( + pt_seg.shape + != ort_seg.shape + ): + raise RuntimeError( + f"shape seg divergiu: PT={pt_seg.shape} ORT={ort_seg.shape}" + ) + + if ( + pt_probs.shape + != ort_probs.shape + ): + raise RuntimeError( + f"shape probs divergiu: PT={pt_probs.shape} ORT={ort_probs.shape}" + ) + + seg_agreement = float( + np.mean( + pt_seg + == ort_seg + ) + ) + + probs_abs = np.abs( + pt_probs.astype( + np.float64 + ) + - ort_probs.astype( + np.float64 + ) + ) + + max_abs = float( + probs_abs.max() + ) + + mean_abs = float( + probs_abs.mean() + ) + + pt_label = int( + np.argmax( + pt_probs[ + 0 + ] + ) + ) + ort_label = int( + np.argmax( + ort_probs[ + 0 + ] + ) + ) + + return { + "seg_agreement": seg_agreement, + "label_same": ( + pt_label + == ort_label + ), + "pt_label": pt_label, + "ort_label": ort_label, + "label_probs_max_abs": max_abs, + "label_probs_mean_abs": mean_abs, + } + + +def verify_with_ort( + *, + onnx_path: Path, + wrapper: CorridorFieldOnnx, + contract: CheckpointContract, + project_root: Path, + device: torch.device, + provider_mode: str, + verify_samples: int, + verify_dir: Optional[Path], + min_seg_agreement: float, + max_prob_abs: float, + max_prob_abs_synthetic: float, +) -> dict: + session, provider_info = create_ort_session( + onnx_path, + provider_mode, + ) + + if session is None: + print( + "[WARN] onnxruntime não instalado; " + "verificação PT x ORT pulada." + ) + + return { + "provider": provider_info, + "executed": False, + "cases": [], + "passed": None, + } + + print( + f"[ORT] providers ativos: " + f"{provider_info['active']}" + ) + + ort_input = session.get_inputs()[ + 0 + ] + + ort_outputs = [ + x.name + for x in session.get_outputs() + ] + + if ( + ort_input.name + != INPUT_NAME + ): + raise RuntimeError( + f"ORT input inesperado: {ort_input.name}" + ) + + if ort_outputs != [ + SEG_OUTPUT_NAME, + STATUS_OUTPUT_NAME, + ]: + raise RuntimeError( + f"ORT outputs inesperados: {ort_outputs}" + ) + + W, H = contract.resolution_wh + + cases = [] + + # Caso sintético. + rng = np.random.default_rng( + 20260914 + ) + + synthetic = rng.integers( + 0, + 256, + size=( + 1, + H, + W, + 3, + ), + dtype=np.uint8, + ) + + source_cases = [ + ( + "synthetic_random", + synthetic, + ) + ] + + images = find_verification_images( + project_root=project_root, + contract=contract, + explicit_dir=verify_dir, + count=verify_samples, + ) + + for image_path in images: + source_cases.append( + ( + str( + image_path + ), + image_to_contract_input( + image_path, + contract, + ), + ) + ) + + # Warmup ORT/CUDA fora das métricas. + # + # A primeira session.run() pode pagar criação de kernels, alocações e + # otimizações do CUDA EP. Isso não representa steady-state do modelo. + # Usamos o caso sintético apenas para aquecer e descartamos esse tempo. + if source_cases: + warmup_input = source_cases[0][1] + + for _ in range(2): + session.run( + [ + SEG_OUTPUT_NAME, + STATUS_OUTPUT_NAME, + ], + { + INPUT_NAME: warmup_input + }, + ) + + print( + "[ORT] warmup concluído (2 inferências não contabilizadas)." + ) + + all_passed = True + + for name, input_np in source_cases: + pt_seg, pt_probs = run_pytorch_wrapper( + wrapper, + input_np, + device, + ) + + t0 = time.perf_counter() + + ort_seg, ort_probs = session.run( + [ + SEG_OUTPUT_NAME, + STATUS_OUTPUT_NAME, + ], + { + INPUT_NAME: input_np + }, + ) + + ort_ms = ( + time.perf_counter() + - t0 + ) * 1000.0 + + cmp = compare_outputs( + pt_seg, + pt_probs, + ort_seg, + ort_probs, + ) + + is_synthetic = ( + name + == "synthetic_random" + ) + + prob_threshold = float( + max_prob_abs_synthetic + if is_synthetic + else max_prob_abs + ) + + passed = ( + cmp[ + "seg_agreement" + ] + >= float( + min_seg_agreement + ) + and cmp[ + "label_same" + ] + and cmp[ + "label_probs_max_abs" + ] + <= prob_threshold + ) + + all_passed = ( + all_passed + and passed + ) + + row = { + "source": name, + "ort_ms": float( + ort_ms + ), + **cmp, + "case_kind": ( + "synthetic" + if is_synthetic + else "real" + ), + "prob_abs_threshold": float( + prob_threshold + ), + "passed": bool( + passed + ), + } + + cases.append( + row + ) + + print( + f"[VERIFY] {name} | " + f"seg={cmp['seg_agreement'] * 100.0:.5f}% " + f"label={cmp['pt_label']}/{cmp['ort_label']} " + f"pMaxAbs={cmp['label_probs_max_abs']:.7g} " + f"(lim={prob_threshold:.7g}) " + f"ORT={ort_ms:.2f}ms | " + f"{'OK' if passed else 'FAIL'}" + ) + + if not all_passed: + raise RuntimeError( + "Verificação PyTorch x ONNX falhou. " + "Não considere este ONNX pronto para campo." + ) + + return { + "provider": provider_info, + "executed": True, + "cases": cases, + "passed": True, + "thresholds": { + "min_seg_agreement": float( + min_seg_agreement + ), + "max_label_probs_abs_real": float( + max_prob_abs + ), + "max_label_probs_abs_synthetic": float( + max_prob_abs_synthetic + ), + "require_same_status_argmax": True, + }, + } + + +# ============================================================================= +# Optional trtexec parse/build validation +# ============================================================================= + +def find_trtexec( + explicit: Optional[ + str + ], +) -> Optional[ + Path +]: + if explicit: + p = Path( + explicit + ) + + if p.is_file(): + return p.resolve() + + found = shutil.which( + "trtexec" + ) + + if found: + return Path( + found + ).resolve() + + return None + + +# ============================================================================= +# Main +# ============================================================================= + +def main(): + ap = argparse.ArgumentParser( + description=( + "Export ONNX campo v2: RGB uint8 NHWC -> seg_ids + label_probs." + ) + ) + + ap.add_argument( + "--config", + default="config.json", + ) + + ap.add_argument( + "--ckpt", + default=None, + help="Checkpoint .pt explícito.", + ) + + ap.add_argument( + "--checkpoint", + default="operational", + choices=[ + "operational", + "nav", + "status", + "last", + ], + help="Checkpoint automático quando --ckpt não é informado.", + ) + + ap.add_argument( + "--out", + default=None, + help="Destino ONNX. Default fica ao lado do checkpoint.", + ) + + ap.add_argument( + "--device", + default="cuda", + choices=[ + "cuda", + "cpu", + ], + ) + + ap.add_argument( + "--opset", + type=int, + default=DEFAULT_OPSET, + ) + + ap.add_argument( + "--verify-provider", + default="auto", + choices=[ + "auto", + "cuda", + "cpu", + "tensorrt", + ], + help=( + "Provider ORT usado no sanity check. " + "auto prefere CUDA, não TensorRT." + ), + ) + + ap.add_argument( + "--verify-samples", + type=int, + default=3, + help="Quantidade de imagens reais do val além do caso sintético.", + ) + + ap.add_argument( + "--verify-dir", + default=None, + help="Pasta de imagens para verificação opcional.", + ) + + ap.add_argument( + "--min-seg-agreement", + type=float, + default=0.999, + help="Acordo mínimo PT x ONNX na máscara final.", + ) + + ap.add_argument( + "--max-prob-abs", + type=float, + default=2e-4, + help=( + "Erro absoluto máximo aceitável nas probabilidades de status " + "para imagens reais." + ), + ) + + ap.add_argument( + "--max-prob-abs-synthetic", + type=float, + default=1e-3, + help=( + "Erro absoluto máximo no caso synthetic_random. " + "Ruído uniforme 0..255 é OOD e usa tolerância numérica separada." + ), + ) + + ap.add_argument( + "--no-ort-verify", + action="store_true", + ) + + ap.add_argument( + "--trtexec", + default=None, + help="Caminho opcional do trtexec; apenas registrado no contrato.", + ) + + args = ap.parse_args() + + config_path = Path( + args.config + ).resolve() + + if not config_path.is_file(): + raise FileNotFoundError( + f"Config não encontrado: {config_path}" + ) + + config = load_json( + config_path + ) + + project_root = resolve_project_root( + config_path + ) + + if args.ckpt: + checkpoint_path = Path( + args.ckpt + ).resolve() + else: + checkpoint_path = default_checkpoint_path( + project_root, + config, + args.checkpoint, + ) + + ckpt, contract = load_checkpoint_contract( + checkpoint_path, + config, + ) + + use_cuda = ( + args.device + == "cuda" + and torch.cuda.is_available() + ) + + if ( + args.device + == "cuda" + and not use_cuda + ): + print( + "[WARN] CUDA solicitado, mas indisponível. " + "Exportando em CPU." + ) + + device = torch.device( + "cuda" + if use_cuda + else "cpu" + ) + + if args.out: + out_path = Path( + args.out + ).resolve() + else: + stem = checkpoint_path.stem + + out_path = ( + checkpoint_path.parent + / f"{stem}.field_v2.onnx" + ) + + contract_path = out_path.with_suffix( + ".contract.json" + ) + + checkpoint_sha = sha256_file( + checkpoint_path + ) + + W, H = contract.resolution_wh + + print("=" * 96) + print( + "AGROBOT | EXPORT ONNX CAMPO V2" + ) + print("=" * 96) + print( + f"Config : {config_path}" + ) + print( + f"Checkpoint : {checkpoint_path}" + ) + print( + f"Checkpoint SHA256 : {checkpoint_sha}" + ) + print( + f"Trainer : {contract.trainer_version}" + ) + print( + f"Epoch : {contract.epoch}" + ) + print( + f"Bests : {contract.bests}" + ) + print( + f"Backbone : {contract.backbone}" + ) + print( + f"Resolution : {W}x{H}" + ) + print( + f"Seg classes : {contract.seg_id2label}" + ) + print( + f"Nav / NonNav : " + f"{contract.nav_class_id} / {contract.non_nav_class_id}" + ) + print( + f"Status classes : {contract.label_name_by_id}" + ) + print( + f"Status head : {contract.status_head}" + ) + print( + f"Norm mean : {contract.norm_mean}" + ) + print( + f"Norm std : {contract.norm_std}" + ) + print( + f"Device export : {device}" + ) + print( + f"Opset : {args.opset}" + ) + print( + f"ONNX : {out_path}" + ) + print("=" * 96) + + print( + "[MODEL] reconstruindo SegFormer + CorridorStatusHead..." + ) + + base_model, status_head = build_models( + ckpt, + contract, + device, + ) + + # Wrapper de REFERÊNCIA: arquitetura exatamente igual ao treino. + reference_wrapper = CorridorFieldOnnx( + base_model=base_model, + status_head=status_head, + resolution_wh=contract.resolution_wh, + norm_mean=contract.norm_mean, + norm_std=contract.norm_std, + ).to( + device + ).eval() + + # ------------------------------------------------------------------------- + # Adaptação export-only do AdaptiveAvgPool2d. + # + # O ONNX de campo possui forma espacial fixa. Descobrimos as formas nativas + # do checkpoint e substituímos SOMENTE NA CÓPIA da cabeça: + # + # AdaptiveAvgPool2d(pool_hw) + # + # por: + # + # AvgPool2d(kernel=input_hw/pool_hw, stride=kernel) + # + # quando a divisão é exata. O wrapper original permanece intacto e será + # usado depois para verificar PyTorch-original x ONNX. + # ------------------------------------------------------------------------- + + feat_hw, seg_hw = discover_status_native_shapes( + base_model=base_model, + H=H, + W=W, + device=device, + mean=contract.norm_mean, + std=contract.norm_std, + ) + + export_status_head, export_pool_plan = build_export_status_head( + reference_status_head=status_head, + feat_hw=feat_hw, + seg_hw=seg_hw, + device=device, + ) + + print( + "[EXPORT-ADAPT] status pooling | " + f"feat {export_pool_plan.feat_hw} -> {export_pool_plan.pool_hw} " + f"kernel={export_pool_plan.feat_kernel_hw} | " + f"seg {export_pool_plan.seg_hw} -> {export_pool_plan.pool_hw} " + f"kernel={export_pool_plan.seg_kernel_hw}" + ) + + export_wrapper = CorridorFieldOnnx( + base_model=base_model, + status_head=export_status_head, + resolution_wh=contract.resolution_wh, + norm_mean=contract.norm_mean, + norm_std=contract.norm_std, + ).to( + device + ).eval() + + export_adaptation_check = verify_export_wrapper_equivalence( + reference_wrapper=reference_wrapper, + export_wrapper=export_wrapper, + H=H, + W=W, + device=device, + max_prob_abs=1e-6, + ) + + # ------------------------------------------------------------------------- + # Export + # ------------------------------------------------------------------------- + + export_onnx( + wrapper=export_wrapper, + out_path=out_path, + H=H, + W=W, + device=device, + opset=int( + args.opset + ), + ) + + print( + f"[OK] ONNX exportado: {out_path}" + ) + + # ------------------------------------------------------------------------- + # Checker + # ------------------------------------------------------------------------- + + onnx_check = check_onnx( + out_path + ) + + # ------------------------------------------------------------------------- + # Contract + metadata embed + # ------------------------------------------------------------------------- + + core_contract = contract_core_dict( + checkpoint_contract=contract, + checkpoint_sha256=checkpoint_sha, + opset=int( + args.opset + ), + export_pool_plan=export_pool_plan, + ) + + metadata_embedded = embed_onnx_metadata( + out_path, + core_contract, + ) + + if metadata_embedded: + print( + "[OK] contrato básico embutido em ONNX metadata_props." + ) + else: + print( + "[WARN] metadata não embutida porque pacote onnx não está instalado." + ) + + # Embedding modifies the file, so checker once more when possible. + if metadata_embedded: + onnx_check = check_onnx( + out_path + ) + + # ------------------------------------------------------------------------- + # ORT verification + # ------------------------------------------------------------------------- + + verify_result = { + "executed": False, + "passed": None, + } + + if not args.no_ort_verify: + verify_dir = ( + Path( + args.verify_dir + ).resolve() + if args.verify_dir + else None + ) + + verify_result = verify_with_ort( + onnx_path=out_path, + wrapper=reference_wrapper, + contract=contract, + project_root=project_root, + device=device, + provider_mode=args.verify_provider, + verify_samples=max( + 0, + int( + args.verify_samples + ), + ), + verify_dir=verify_dir, + min_seg_agreement=float( + args.min_seg_agreement + ), + max_prob_abs=float( + args.max_prob_abs + ), + max_prob_abs_synthetic=float( + args.max_prob_abs_synthetic + ), + ) + + # ------------------------------------------------------------------------- + # Final hashes / full sidecar + # ------------------------------------------------------------------------- + + onnx_sha = sha256_file( + out_path + ) + + trtexec_path = find_trtexec( + args.trtexec + ) + + full_contract = { + **core_contract, + + "artifact": { + "onnx_path": str( + out_path + ), + "onnx_sha256": onnx_sha, + "contract_path": str( + contract_path + ), + "checkpoint_path": str( + checkpoint_path + ), + "checkpoint_sha256": checkpoint_sha, + "metadata_embedded_in_onnx": bool( + metadata_embedded + ), + }, + + "validation": { + "export_adaptation_pytorch_equivalence": export_adaptation_check, + "onnx_check": onnx_check, + "ort": verify_result, + }, + + "deployment": { + "target": "VisualWorker / ONNX Runtime / TensorRT EP", + "recommended_precision": "FP16 TensorRT engine", + "trtexec_detected": ( + str( + trtexec_path + ) + if trtexec_path + is not None + else None + ), + "runtime_steps": [ + "obtain RGB uint8 frame", + ( + f"crop source ROI [{contract.roi_inicio}, " + f"{contract.roi_tamanho}] if source frame is larger" + ), + ( + f"resize source image to {W}x{H} " + "before ONNX input" + ), + ( + f"feed {INPUT_NAME} as contiguous uint8 NHWC " + f"shape [1,{H},{W},3]" + ), + f"read {SEG_OUTPUT_NAME} int32 [1,{H},{W}]", + ( + f"read {STATUS_OUTPUT_NAME} float32 " + f"[1,{len(contract.label_name_by_id)}]" + ), + "status_id = argmax(label_probs[0])", + "status_conf = label_probs[0,status_id]", + ], + }, + } + + save_json( + contract_path, + full_contract, + ) + + print() + print("=" * 96) + print( + "EXPORT FINALIZADO" + ) + print("=" * 96) + print( + f"ONNX : {out_path}" + ) + print( + f"ONNX SHA256 : {onnx_sha}" + ) + print( + f"Contrato : {contract_path}" + ) + print( + f"Input : " + f"{INPUT_NAME} uint8 [1,{H},{W},3] RGB NHWC" + ) + print( + f"Output SEG : " + f"{SEG_OUTPUT_NAME} int32 [1,{H},{W}]" + ) + print( + f"Output STATUS : " + f"{STATUS_OUTPUT_NAME} float32 " + f"[1,{len(contract.label_name_by_id)}]" + ) + print( + f"ORT verify : " + f"{verify_result.get('passed')}" + ) + print("=" * 96) + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/oak-d/_7_benchmark_corridor.py b/Python/OAK/datasets/oak-d/_7_benchmark_corridor.py new file mode 100644 index 000000000..e60733d3b --- /dev/null +++ b/Python/OAK/datasets/oak-d/_7_benchmark_corridor.py @@ -0,0 +1,3589 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +benchmark_corridor_field_v2.py +============================== + +Benchmark de ponta a ponta do modelo frontal de corredor. + +Fecha o pipeline: + + CAPTURE + ↓ + ANNOTATE + ↓ + NORMALIZE + ↓ + SPLIT + ↓ + TRAIN + ↓ + TEST / REVIEW + ↓ + EXPORT ONNX + ↓ + BENCHMARK ← este script + ↓ + VISUAL WORKER + +O benchmark foi dividido em camadas para não esconder gargalos. + +FASES OFFLINE +------------- +1) host_prepare + Simula o preparo do frame bruto: + BGR source + ROI vertical + resize INTER_AREA + BGR -> RGB + contiguous uint8 NHWC + batch [1,H,W,3] + +2) pytorch_core + Mede SegFormer + CorridorStatusHead + resize logits + argmax + softmax. + Entrada já está normalizada float32 NCHW. + É a melhor aproximação do "modelo puro". + +3) pytorch_field + Mede o wrapper de campo PyTorch: + uint8 NHWC + cast + /255 + transpose + mean/std + modelo + resize + argmax + status softmax + +4) onnx_cuda + Mede session.run() do ONNX de campo com CUDA EP. + Inclui o contrato real: + uint8 host -> GPU -> outputs CPU + +5) onnx_tensorrt + Mede o MESMO ONNX com TensorRT EP. + Registra: + criação da sessão + primeira inferência + warmup + steady-state + +PIPELINE REAL COM OAK-D +----------------------- +Opcional com --camera. + +Threads independentes: + + CaptureThread + OAK-D Lite BGR 1080p + ↓ latest-slot + PrepareThread + ROI + resize + BGR→RGB + contiguous + ↓ latest-slot + InferThread + ONNX Runtime provider escolhido + ↓ + seg_ids + label_probs + +A política latest-frame é intencional: + se o consumidor atrasar, frames antigos são descartados. + Isso evita que a inferência processe uma fila envelhecida. + +Métricas de câmera: + capture_fps + prepare_fps + output_fps + + capture_frames + prepared_frames + inferred_frames + + dropped_before_prepare + dropped_before_infer + + prepare_ms: + mean / p50 / p95 / p99 / max + + inference_ms: + mean / p50 / p95 / p99 / max + + age_before_infer_ms + wait_after_prepare_ms + capture_to_result_ms + +ARTEFATOS +--------- +benchmark_corridor// + benchmark_summary.json + benchmark_results.csv + camera_pipeline_samples.csv # se --camera + environment.json + +USO RÁPIDO +---------- +Benchmark offline completo: + + python benchmark_corridor_field_v2.py + +Adicionar pipeline real da OAK-D com TensorRT: + + python benchmark_corridor_field_v2.py --camera + +Pipeline real usando CUDA EP: + + python benchmark_corridor_field_v2.py ^ + --camera ^ + --pipeline_provider cuda + +Medir cold-start do TensorRT: + + python benchmark_corridor_field_v2.py ^ + --clear_trt_cache + +Rodar mais tempo: + + python benchmark_corridor_field_v2.py ^ + --iterations 500 ^ + --warmup 50 ^ + --camera ^ + --camera_duration 60 + +Somente fases específicas: + + python benchmark_corridor_field_v2.py ^ + --modes pytorch_core,onnx_tensorrt + +OBSERVAÇÕES +----------- +- O script usa test_corridor_agri_v2.py para reconstruir o checkpoint. +- Usa export_corridor_agri_onnx_v2.py para o wrapper de campo PyTorch. +- O ONNX default é: + .field_v2.onnx +- O benchmark não compila o ONNX. Rode o export primeiro. +""" + +from __future__ import annotations + +import argparse +import csv +import gc +import importlib.util +import json +import os +import platform +import queue +import shutil +import socket +import statistics +import sys +import threading +import time +from dataclasses import dataclass, asdict +from datetime import datetime +from pathlib import Path +from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple + +import cv2 +import numpy as np +import torch +import torch.nn.functional as F + + +# ============================================================================= +# Constants +# ============================================================================= + +BENCHMARK_VERSION = "agrobot_corridor_field_benchmark_v2.0" + +DEFAULT_MODES = ( + "host_prepare", + "pytorch_core", + "pytorch_field", + "onnx_cuda", + "onnx_tensorrt", +) + +VALID_MODES = { + "host_prepare", + "pytorch_core", + "pytorch_field", + "onnx_cpu", + "onnx_cuda", + "onnx_tensorrt", +} + +INPUT_NAME = "rgb_u8_nhwc" +SEG_OUTPUT_NAME = "seg_ids" +STATUS_OUTPUT_NAME = "label_probs" + + +# ============================================================================= +# Generic helpers +# ============================================================================= + +def now_iso() -> str: + return datetime.now().isoformat( + timespec="seconds" + ) + + +def timestamp_run_name() -> str: + return datetime.now().strftime( + "%Y%m%d_%H%M%S" + ) + + +def safe_json_load( + path: Path, +) -> dict: + with path.open( + "r", + encoding="utf-8", + ) as f: + data = json.load( + f + ) + + if not isinstance( + data, + dict, + ): + raise ValueError( + f"Esperava objeto JSON: {path}" + ) + + return data + + +def save_json( + path: Path, + data: dict, +) -> None: + path.parent.mkdir( + parents=True, + exist_ok=True, + ) + + tmp = path.with_suffix( + path.suffix + + ".tmp" + ) + + with tmp.open( + "w", + encoding="utf-8", + ) as f: + json.dump( + data, + f, + ensure_ascii=False, + indent=2, + ) + + os.replace( + tmp, + path, + ) + + +def write_csv( + path: Path, + rows: Sequence[dict], +) -> None: + path.parent.mkdir( + parents=True, + exist_ok=True, + ) + + if not rows: + path.write_text( + "", + encoding="utf-8", + ) + return + + fields = [] + seen = set() + + for row in rows: + for key in row: + if key not in seen: + seen.add( + key + ) + fields.append( + key + ) + + with path.open( + "w", + newline="", + encoding="utf-8-sig", + ) as f: + writer = csv.DictWriter( + f, + fieldnames=fields, + extrasaction="ignore", + ) + writer.writeheader() + writer.writerows( + rows + ) + + +def load_module_from_file( + script_path: Path, + module_name: str, +): + if not script_path.is_file(): + raise FileNotFoundError( + f"Script dependência não encontrado: {script_path}" + ) + + spec = importlib.util.spec_from_file_location( + module_name, + str( + script_path + ), + ) + + if ( + spec is None + or spec.loader is None + ): + raise RuntimeError( + f"Não consegui importar {script_path}" + ) + + module = importlib.util.module_from_spec( + spec + ) + + sys.modules[ + module_name + ] = module + + spec.loader.exec_module( + module + ) + + return module + + +def resolve_project_root( + config_path: Path, +) -> Path: + root = config_path.resolve().parent + + if not ( + root + / "dataset" + ).is_dir(): + raise FileNotFoundError( + f"dataset/ não encontrado ao lado de {config_path}" + ) + + return root + + +def parse_modes( + text: str, +) -> List[str]: + result = [] + + for raw in str( + text + ).split( + "," + ): + mode = raw.strip().lower() + + if not mode: + continue + + if mode == "all": + for item in DEFAULT_MODES: + if item not in result: + result.append( + item + ) + continue + + if mode not in VALID_MODES: + raise ValueError( + f"Modo desconhecido: {mode}. " + f"Válidos={sorted(VALID_MODES)}" + ) + + if mode not in result: + result.append( + mode + ) + + if not result: + raise ValueError( + "Nenhum modo selecionado." + ) + + return result + + +def percentile( + values: Sequence[float], + q: float, +) -> Optional[float]: + if not values: + return None + + return float( + np.percentile( + np.asarray( + values, + dtype=np.float64, + ), + q, + ) + ) + + +def stats_ms( + values: Sequence[float], +) -> dict: + if not values: + return { + "n": 0, + "min": None, + "mean": None, + "p50": None, + "p95": None, + "p99": None, + "max": None, + "std": None, + } + + arr = np.asarray( + values, + dtype=np.float64, + ) + + return { + "n": int( + arr.size + ), + "min": float( + arr.min() + ), + "mean": float( + arr.mean() + ), + "p50": float( + np.percentile( + arr, + 50, + ) + ), + "p95": float( + np.percentile( + arr, + 95, + ) + ), + "p99": float( + np.percentile( + arr, + 99, + ) + ), + "max": float( + arr.max() + ), + "std": float( + arr.std() + ), + } + + +def fps_from_ms( + mean_ms: Optional[float], +) -> Optional[float]: + if ( + mean_ms is None + or mean_ms <= 0 + ): + return None + + return float( + 1000.0 + / mean_ms + ) + + +def flatten_stats_row( + name: str, + kind: str, + stat: dict, + extra: Optional[dict] = None, +) -> dict: + row = { + "name": name, + "kind": kind, + "n": stat.get( + "n" + ), + "mean_ms": stat.get( + "mean" + ), + "p50_ms": stat.get( + "p50" + ), + "p95_ms": stat.get( + "p95" + ), + "p99_ms": stat.get( + "p99" + ), + "min_ms": stat.get( + "min" + ), + "max_ms": stat.get( + "max" + ), + "std_ms": stat.get( + "std" + ), + "fps_from_mean": fps_from_ms( + stat.get( + "mean" + ) + ), + } + + if extra: + row.update( + extra + ) + + return row + + +def synchronize_cuda( + device: torch.device, +) -> None: + if device.type == "cuda": + torch.cuda.synchronize() + + +# ============================================================================= +# Environment +# ============================================================================= + +def collect_environment() -> dict: + env = { + "created_at": now_iso(), + "benchmark_version": BENCHMARK_VERSION, + "hostname": socket.gethostname(), + "platform": platform.platform(), + "python": sys.version, + "opencv": cv2.__version__, + "numpy": np.__version__, + "torch": torch.__version__, + "torch_cuda_available": bool( + torch.cuda.is_available() + ), + "torch_cuda_version": torch.version.cuda, + "cudnn_version": ( + torch.backends.cudnn.version() + if torch.backends.cudnn.is_available() + else None + ), + } + + if torch.cuda.is_available(): + env.update({ + "gpu_name": torch.cuda.get_device_name( + 0 + ), + "gpu_capability": list( + torch.cuda.get_device_capability( + 0 + ) + ), + "gpu_total_memory_bytes": int( + torch.cuda.get_device_properties( + 0 + ).total_memory + ), + }) + + try: + import onnxruntime as ort + + env[ + "onnxruntime" + ] = ort.__version__ + + env[ + "onnxruntime_available_providers" + ] = ort.get_available_providers() + + except Exception as exc: + env[ + "onnxruntime" + ] = None + + env[ + "onnxruntime_error" + ] = ( + f"{type(exc).__name__}: {exc}" + ) + + try: + import depthai as dai + + env[ + "depthai" + ] = getattr( + dai, + "__version__", + "unknown", + ) + except Exception: + env[ + "depthai" + ] = None + + return env + + +# ============================================================================= +# Geometry / preprocess +# ============================================================================= + +def compute_roi( + h: int, + roi_inicio: float, + roi_tamanho: float, +) -> Tuple[int, int]: + y0 = int( + round( + h + * float( + roi_inicio + ) + ) + ) + + y1 = int( + round( + h + * min( + 1.0, + float( + roi_inicio + + roi_tamanho + ), + ) + ) + ) + + y0 = max( + 0, + min( + h - 1, + y0, + ), + ) + + y1 = max( + y0 + 1, + min( + h, + y1, + ), + ) + + return ( + y0, + y1, + ) + + +def prepare_rgb_u8_nhwc( + frame_bgr: np.ndarray, + *, + resolution_wh: Tuple[int, int], + roi_inicio: float, + roi_tamanho: float, +) -> np.ndarray: + """ + Contrato EXTERNO do ONNX v2. + + frame BGR source + -> ROI vertical + -> resize INTER_AREA + -> RGB + -> uint8 contiguous + -> [1,H,W,3] + """ + if ( + frame_bgr.ndim + != 3 + or frame_bgr.shape[ + 2 + ] + < 3 + ): + raise ValueError( + f"Frame BGR inválido: shape={frame_bgr.shape}" + ) + + h, w = frame_bgr.shape[ + :2 + ] + + y0, y1 = compute_roi( + h, + roi_inicio, + roi_tamanho, + ) + + roi = frame_bgr[ + y0:y1, + :, + :3, + ] + + W, H = ( + int( + resolution_wh[ + 0 + ] + ), + int( + resolution_wh[ + 1 + ] + ), + ) + + if roi.shape[ + 1 + ] == W and roi.shape[ + 0 + ] == H: + resized = roi + else: + resized = cv2.resize( + roi, + ( + W, + H, + ), + interpolation=cv2.INTER_AREA, + ) + + rgb = cv2.cvtColor( + resized, + cv2.COLOR_BGR2RGB, + ) + + return np.ascontiguousarray( + rgb[ + None, + ..., + ], + dtype=np.uint8, + ) + + +def make_synthetic_source_frame( + source_wh: Tuple[int, int], + seed: int = 20260914, +) -> np.ndarray: + W, H = ( + int( + source_wh[ + 0 + ] + ), + int( + source_wh[ + 1 + ] + ), + ) + + rng = np.random.default_rng( + seed + ) + + return rng.integers( + 0, + 256, + size=( + H, + W, + 3, + ), + dtype=np.uint8, + ) + + +def find_reference_image( + project_root: Path, + resolution_wh: Tuple[int, int], +) -> Optional[Path]: + W, H = resolution_wh + + candidates = [ + ( + project_root + / "dataset" + / f"{W}x{H}" + / "split" + / "val" + / "images" + ), + ( + project_root + / "dataset" + / f"{W}x{H}" + / "split" + / "train" + / "images" + ), + ] + + exts = { + ".png", + ".jpg", + ".jpeg", + ".bmp", + ".webp", + } + + for folder in candidates: + if not folder.is_dir(): + continue + + images = sorted( + [ + p + for p in folder.iterdir() + if ( + p.is_file() + and p.suffix.lower() + in exts + ) + ], + key=lambda p: p.name.lower(), + ) + + if images: + return images[ + 0 + ] + + return None + + +# ============================================================================= +# PyTorch inputs and calls +# ============================================================================= + +def uint8_contract_to_normalized_nchw( + input_np: np.ndarray, + *, + mean: Sequence[float], + std: Sequence[float], + device: torch.device, +) -> torch.Tensor: + x = torch.from_numpy( + np.ascontiguousarray( + input_np + ) + ).to( + device, + non_blocking=True, + ) + + x = x.to( + torch.float32 + ) + + x = ( + x + / 255.0 + ) + + x = x.permute( + 0, + 3, + 1, + 2, + ) + + mean_t = torch.tensor( + list( + mean + ), + dtype=torch.float32, + device=device, + ).view( + 1, + 3, + 1, + 1, + ) + + std_t = torch.tensor( + list( + std + ), + dtype=torch.float32, + device=device, + ).view( + 1, + 3, + 1, + 1, + ) + + return ( + x + - mean_t + ) / std_t + + +@torch.inference_mode() +def pytorch_core_call( + base_model, + status_head, + x_norm: torch.Tensor, + resolution_wh: Tuple[int, int], +): + out = base_model( + pixel_values=x_norm, + output_hidden_states=True, + return_dict=True, + ) + + seg_logits_native = out.logits + feat = out.hidden_states[ + -1 + ] + + status_logits = status_head( + feat, + seg_logits_native, + ) + + W, H = resolution_wh + + seg_logits_full = F.interpolate( + seg_logits_native, + size=( + H, + W, + ), + mode="bilinear", + align_corners=False, + ) + + seg_ids = torch.argmax( + seg_logits_full, + dim=1, + ) + + label_probs = torch.softmax( + status_logits, + dim=1, + ) + + return ( + seg_ids, + label_probs, + ) + + +# ============================================================================= +# Microbench harness +# ============================================================================= + +def benchmark_callable( + *, + name: str, + call, + warmup: int, + iterations: int, + sync_before_after: Optional[ + callable + ] = None, +) -> dict: + """ + Mede wall time por chamada. + + Para PyTorch CUDA: + sync_before_after = torch.cuda.synchronize + + Para ORT session.run: + chamada é síncrona do ponto de vista host, portanto não precisa. + """ + print( + f"[BENCH] {name}: warmup={warmup} iterations={iterations}" + ) + + first_ms = None + + if sync_before_after: + sync_before_after() + + t0 = time.perf_counter() + + call() + + if sync_before_after: + sync_before_after() + + first_ms = ( + time.perf_counter() + - t0 + ) * 1000.0 + + for _ in range( + max( + 0, + int( + warmup + ), + ) + ): + call() + + if sync_before_after: + sync_before_after() + + samples = [] + + for _ in range( + max( + 1, + int( + iterations + ), + ) + ): + if sync_before_after: + sync_before_after() + + t0 = time.perf_counter() + + call() + + if sync_before_after: + sync_before_after() + + samples.append( + ( + time.perf_counter() + - t0 + ) + * 1000.0 + ) + + stat = stats_ms( + samples + ) + + print( + f" first={first_ms:.3f}ms | " + f"mean={stat['mean']:.3f}ms | " + f"p50={stat['p50']:.3f} | " + f"p95={stat['p95']:.3f} | " + f"p99={stat['p99']:.3f} | " + f"fps≈{fps_from_ms(stat['mean']):.2f}" + ) + + return { + "name": name, + "first_ms": float( + first_ms + ), + "latency_ms": stat, + "fps_from_mean": fps_from_ms( + stat[ + "mean" + ] + ), + } + + +# ============================================================================= +# ONNX Runtime +# ============================================================================= + +def ort_provider_name( + provider_mode: str, +) -> str: + return { + "cpu": "CPUExecutionProvider", + "cuda": "CUDAExecutionProvider", + "tensorrt": "TensorrtExecutionProvider", + }[ + provider_mode + ] + + +def create_ort_session( + *, + onnx_path: Path, + provider_mode: str, + trt_cache_dir: Path, + trt_fp16: bool, +) -> Tuple[ + Any, + dict, +]: + import onnxruntime as ort + + available = ort.get_available_providers() + + provider_mode = str( + provider_mode + ).lower() + + if provider_mode == "cpu": + requested = [ + "CPUExecutionProvider" + ] + + elif provider_mode == "cuda": + requested = [ + "CUDAExecutionProvider", + "CPUExecutionProvider", + ] + + elif provider_mode == "tensorrt": + trt_cache_dir.mkdir( + parents=True, + exist_ok=True, + ) + + trt_options = { + "trt_engine_cache_enable": True, + "trt_engine_cache_path": str( + trt_cache_dir + ), + "trt_timing_cache_enable": True, + "trt_timing_cache_path": str( + trt_cache_dir + ), + "trt_fp16_enable": bool( + trt_fp16 + ), + } + + requested = [ + ( + "TensorrtExecutionProvider", + trt_options, + ), + "CUDAExecutionProvider", + "CPUExecutionProvider", + ] + + else: + raise ValueError( + provider_mode + ) + + filtered = [] + + for provider in requested: + name = ( + provider[ + 0 + ] + if isinstance( + provider, + tuple, + ) + else provider + ) + + if name in available: + filtered.append( + provider + ) + + target_name = ort_provider_name( + provider_mode + ) + + filtered_names = [ + p[ + 0 + ] + if isinstance( + p, + tuple, + ) + else p + for p in filtered + ] + + if target_name not in filtered_names: + raise RuntimeError( + f"Provider {target_name} indisponível. " + f"ORT disponíveis={available}" + ) + + options = ort.SessionOptions() + + options.graph_optimization_level = ( + ort.GraphOptimizationLevel + .ORT_ENABLE_ALL + ) + + t0 = time.perf_counter() + + session = ort.InferenceSession( + str( + onnx_path + ), + sess_options=options, + providers=filtered, + ) + + build_s = ( + time.perf_counter() + - t0 + ) + + active = session.get_providers() + + if target_name not in active: + raise RuntimeError( + f"Provider solicitado não ficou ativo. " + f"requested={provider_mode} active={active}" + ) + + info = { + "provider_mode": provider_mode, + "target_provider": target_name, + "available_providers": list( + available + ), + "active_providers": list( + active + ), + "session_create_seconds": float( + build_s + ), + "trt_cache_dir": ( + str( + trt_cache_dir + ) + if provider_mode + == "tensorrt" + else None + ), + "trt_fp16": ( + bool( + trt_fp16 + ) + if provider_mode + == "tensorrt" + else None + ), + } + + return ( + session, + info, + ) + + +def benchmark_ort( + *, + onnx_path: Path, + provider_mode: str, + input_np: np.ndarray, + warmup: int, + iterations: int, + trt_cache_dir: Path, + trt_fp16: bool, +) -> dict: + print() + print( + f"[ORT] criando sessão provider={provider_mode}" + ) + + session, session_info = create_ort_session( + onnx_path=onnx_path, + provider_mode=provider_mode, + trt_cache_dir=trt_cache_dir, + trt_fp16=trt_fp16, + ) + + print( + f" sessão={session_info['session_create_seconds']:.3f}s " + f"active={session_info['active_providers']}" + ) + + inputs = session.get_inputs() + + outputs = session.get_outputs() + + if len( + inputs + ) != 1: + raise RuntimeError( + f"ONNX esperado 1 input, recebeu {[x.name for x in inputs]}" + ) + + if inputs[ + 0 + ].name != INPUT_NAME: + raise RuntimeError( + f"Input ONNX inesperado: {inputs[0].name}" + ) + + output_names = [ + x.name + for x in outputs + ] + + if output_names != [ + SEG_OUTPUT_NAME, + STATUS_OUTPUT_NAME, + ]: + raise RuntimeError( + f"Outputs ONNX inesperados: {output_names}" + ) + + def call(): + return session.run( + [ + SEG_OUTPUT_NAME, + STATUS_OUTPUT_NAME, + ], + { + INPUT_NAME: input_np + }, + ) + + result = benchmark_callable( + name=f"onnx_{provider_mode}", + call=call, + warmup=warmup, + iterations=iterations, + sync_before_after=None, + ) + + # Um sanity simples dos outputs. + seg, probs = call() + + result.update({ + "provider": session_info, + "seg_shape": list( + seg.shape + ), + "seg_dtype": str( + seg.dtype + ), + "status_shape": list( + probs.shape + ), + "status_dtype": str( + probs.dtype + ), + "status_argmax": int( + np.argmax( + probs[ + 0 + ] + ) + ), + "status_confidence": float( + np.max( + probs[ + 0 + ] + ) + ), + }) + + return ( + result, + session, + ) + + +# ============================================================================= +# Camera pipeline data +# ============================================================================= + +@dataclass +class CapturedFrame: + frame_id: int + t_capture: float + frame_bgr: np.ndarray + + +@dataclass +class PreparedFrame: + prep_id: int + source_frame_id: int + + t_capture: float + t_prepare_start: float + t_prepare_end: float + + input_np: np.ndarray + + +class LatestSlot: + """ + Slot single-item com versionamento. + + put() substitui sempre o item anterior. + O consumidor pede get_after(version), então nunca espera por frames antigos. + """ + + def __init__( + self, + ): + self._lock = threading.Lock() + self._cond = threading.Condition( + self._lock + ) + + self._version = 0 + self._item = None + self._closed = False + + def put( + self, + item, + ) -> int: + with self._cond: + if self._closed: + return self._version + + self._version += 1 + self._item = item + + self._cond.notify_all() + + return self._version + + def get_after( + self, + last_version: int, + timeout: float = 0.5, + ): + deadline = ( + time.perf_counter() + + float( + timeout + ) + ) + + with self._cond: + while ( + not self._closed + and self._version + <= last_version + ): + remaining = ( + deadline + - time.perf_counter() + ) + + if remaining <= 0: + return ( + last_version, + None, + ) + + self._cond.wait( + timeout=remaining + ) + + if ( + self._version + <= last_version + ): + return ( + last_version, + None, + ) + + return ( + self._version, + self._item, + ) + + def close( + self, + ): + with self._cond: + self._closed = True + self._cond.notify_all() + + +# ============================================================================= +# OAK camera +# ============================================================================= + +class OakCamera: + """ + BGR host frames. + + Tenta API moderna e depois clássica. + """ + + def __init__( + self, + fps: float, + resolution: str, + ): + self.fps = float( + fps + ) + + self.resolution = str( + resolution + ).lower() + + self.mode = None + self.pipeline = None + self.device = None + self.queue = None + self.dai = None + + def start( + self, + ): + try: + import depthai as dai + except ImportError as exc: + raise ImportError( + "Modo câmera exige depthai." + ) from exc + + self.dai = dai + + target_size = ( + ( + 1920, + 1080, + ) + if self.resolution + == "1080p" + else ( + 1280, + 720, + ) + ) + + # API moderna. + try: + pipeline = dai.Pipeline() + + cam = pipeline.create( + dai.node.Camera + ).build() + + out = cam.requestOutput( + size=target_size, + type=dai.ImgFrame.Type.BGR888p, + fps=self.fps, + ) + + q = out.createOutputQueue( + maxSize=4, + blocking=False, + ) + + pipeline.start() + + self.mode = "modern" + self.pipeline = pipeline + self.queue = q + + print( + f"[CAM] modern API " + f"{target_size[0]}x{target_size[1]} " + f"@{self.fps:g}" + ) + + return self + + except Exception as exc: + print( + f"[CAM] modern API indisponível: " + f"{type(exc).__name__}: {exc}" + ) + + # API clássica. + pipeline = dai.Pipeline() + + cam = pipeline.create( + dai.node.ColorCamera + ) + + xout = pipeline.create( + dai.node.XLinkOut + ) + + xout.setStreamName( + "rgb" + ) + + try: + cam.setBoardSocket( + dai.CameraBoardSocket.CAM_A + ) + except Exception: + pass + + if ( + self.resolution + == "1080p" + ): + cam.setResolution( + dai.ColorCameraProperties + .SensorResolution + .THE_1080_P + ) + else: + cam.setResolution( + dai.ColorCameraProperties + .SensorResolution + .THE_720_P + ) + + cam.setFps( + self.fps + ) + + cam.setInterleaved( + False + ) + + cam.video.link( + xout.input + ) + + device = dai.Device( + pipeline + ) + + q = device.getOutputQueue( + name="rgb", + maxSize=4, + blocking=False, + ) + + self.mode = "classic" + self.pipeline = pipeline + self.device = device + self.queue = q + + print( + f"[CAM] classic API " + f"{target_size[0]}x{target_size[1]} " + f"@{self.fps:g}" + ) + + return self + + def read( + self, + ) -> np.ndarray: + if self.queue is None: + raise RuntimeError( + "Câmera não inicializada." + ) + + frame = self.queue.get() + + img = frame.getCvFrame() + + if img is None: + raise RuntimeError( + "Frame OAK vazio." + ) + + return img + + def close( + self, + ): + try: + if ( + self.mode + == "modern" + and self.pipeline + is not None + ): + self.pipeline.stop() + except Exception: + pass + + try: + if self.device is not None: + self.device.close() + except Exception: + pass + + +# ============================================================================= +# Threaded camera benchmark +# ============================================================================= + +class CameraPipelineBenchmark: + def __init__( + self, + *, + session, + contract, + camera_fps: float, + camera_resolution: str, + warmup_s: float, + duration_s: float, + ): + self.session = session + self.contract = contract + + self.camera_fps = float( + camera_fps + ) + + self.camera_resolution = str( + camera_resolution + ) + + self.warmup_s = float( + warmup_s + ) + + self.duration_s = float( + duration_s + ) + + self.raw_slot = LatestSlot() + self.prepared_slot = LatestSlot() + + self.stop_event = threading.Event() + + self.capture_count_total = 0 + self.prepare_count_total = 0 + self.infer_count_total = 0 + + self.capture_times_measure = [] + self.prepare_times_measure = [] + self.infer_times_measure = [] + + self.samples = [] + + self.dropped_before_prepare_measure = 0 + self.dropped_before_infer_measure = 0 + + self.errors = [] + + self.measure_start = None + self.measure_end = None + + self.camera = None + + def in_measurement( + self, + t: Optional[ + float + ] = None, + ) -> bool: + if ( + self.measure_start + is None + or self.measure_end + is None + ): + return False + + if t is None: + t = time.perf_counter() + + return ( + self.measure_start + <= t + <= self.measure_end + ) + + def capture_loop( + self, + ): + frame_id = 0 + + try: + while not self.stop_event.is_set(): + frame = self.camera.read() + + t_capture = time.perf_counter() + + item = CapturedFrame( + frame_id=frame_id, + t_capture=t_capture, + frame_bgr=frame, + ) + + self.raw_slot.put( + item + ) + + self.capture_count_total += 1 + + if self.in_measurement( + t_capture + ): + self.capture_times_measure.append( + t_capture + ) + + frame_id += 1 + + except Exception as exc: + self.errors.append( + "CaptureThread: " + f"{type(exc).__name__}: {exc}" + ) + self.stop_event.set() + + finally: + self.raw_slot.close() + + def prepare_loop( + self, + ): + last_version = 0 + prep_id = 0 + last_source_measure_id = None + + try: + while not self.stop_event.is_set(): + version, item = self.raw_slot.get_after( + last_version, + timeout=0.5, + ) + + if item is None: + continue + + last_version = version + + t0 = time.perf_counter() + + input_np = prepare_rgb_u8_nhwc( + item.frame_bgr, + resolution_wh=self.contract.resolution_wh, + roi_inicio=self.contract.roi_inicio, + roi_tamanho=self.contract.roi_tamanho, + ) + + t1 = time.perf_counter() + + prepared = PreparedFrame( + prep_id=prep_id, + source_frame_id=item.frame_id, + t_capture=item.t_capture, + t_prepare_start=t0, + t_prepare_end=t1, + input_np=input_np, + ) + + self.prepared_slot.put( + prepared + ) + + self.prepare_count_total += 1 + + if self.in_measurement( + t1 + ): + if last_source_measure_id is not None: + gap = ( + item.frame_id + - last_source_measure_id + - 1 + ) + + if gap > 0: + self.dropped_before_prepare_measure += int( + gap + ) + + last_source_measure_id = int( + item.frame_id + ) + + self.prepare_times_measure.append( + t1 + ) + + prep_id += 1 + + except Exception as exc: + self.errors.append( + "PrepareThread: " + f"{type(exc).__name__}: {exc}" + ) + self.stop_event.set() + + finally: + self.prepared_slot.close() + + def infer_loop( + self, + ): + last_version = 0 + + prev_prep_measure_id = None + + try: + while not self.stop_event.is_set(): + version, item = self.prepared_slot.get_after( + last_version, + timeout=0.5, + ) + + if item is None: + continue + + last_version = version + + t_infer_start = time.perf_counter() + + seg_ids, label_probs = self.session.run( + [ + SEG_OUTPUT_NAME, + STATUS_OUTPUT_NAME, + ], + { + INPUT_NAME: item.input_np + }, + ) + + # Visual Worker ainda precisa apenas desta decisão minúscula. + status_id = int( + np.argmax( + label_probs[ + 0 + ] + ) + ) + + status_conf = float( + label_probs[ + 0, + status_id, + ] + ) + + t_done = time.perf_counter() + + self.infer_count_total += 1 + + if self.in_measurement( + t_done + ): + if prev_prep_measure_id is not None: + gap = ( + item.prep_id + - prev_prep_measure_id + - 1 + ) + + if gap > 0: + self.dropped_before_infer_measure += int( + gap + ) + + prev_prep_measure_id = int( + item.prep_id + ) + + self.infer_times_measure.append( + t_done + ) + + self.samples.append({ + "source_frame_id": int( + item.source_frame_id + ), + "prep_id": int( + item.prep_id + ), + + "prepare_ms": float( + ( + item.t_prepare_end + - item.t_prepare_start + ) + * 1000.0 + ), + + "age_at_prepare_start_ms": float( + ( + item.t_prepare_start + - item.t_capture + ) + * 1000.0 + ), + + "wait_after_prepare_ms": float( + ( + t_infer_start + - item.t_prepare_end + ) + * 1000.0 + ), + + "age_before_infer_ms": float( + ( + t_infer_start + - item.t_capture + ) + * 1000.0 + ), + + "inference_ms": float( + ( + t_done + - t_infer_start + ) + * 1000.0 + ), + + "capture_to_result_ms": float( + ( + t_done + - item.t_capture + ) + * 1000.0 + ), + + "status_id": status_id, + "status_confidence": status_conf, + + # Snapshots acumulados. + "dropped_before_prepare_accum": int( + self.dropped_before_prepare_measure + ), + "dropped_before_infer_accum": int( + self.dropped_before_infer_measure + ), + }) + + except Exception as exc: + self.errors.append( + "InferThread: " + f"{type(exc).__name__}: {exc}" + ) + self.stop_event.set() + + def run( + self, + ) -> dict: + self.camera = OakCamera( + fps=self.camera_fps, + resolution=self.camera_resolution, + ).start() + + capture_thread = threading.Thread( + target=self.capture_loop, + name="CaptureThread", + daemon=True, + ) + + prepare_thread = threading.Thread( + target=self.prepare_loop, + name="PrepareThread", + daemon=True, + ) + + infer_thread = threading.Thread( + target=self.infer_loop, + name="InferThread", + daemon=True, + ) + + t_pipeline_start = time.perf_counter() + + capture_thread.start() + prepare_thread.start() + infer_thread.start() + + print( + f"[CAM BENCH] warmup={self.warmup_s:.1f}s" + ) + + while ( + time.perf_counter() + - t_pipeline_start + < self.warmup_s + ): + if self.stop_event.is_set(): + break + + time.sleep( + 0.05 + ) + + self.measure_start = time.perf_counter() + self.measure_end = ( + self.measure_start + + self.duration_s + ) + + print( + f"[CAM BENCH] medição={self.duration_s:.1f}s" + ) + + while ( + time.perf_counter() + < self.measure_end + ): + if self.stop_event.is_set(): + break + + time.sleep( + 0.05 + ) + + self.stop_event.set() + + self.raw_slot.close() + self.prepared_slot.close() + + capture_thread.join( + timeout=3.0 + ) + prepare_thread.join( + timeout=3.0 + ) + infer_thread.join( + timeout=3.0 + ) + + self.camera.close() + + actual_measure_duration = max( + 1e-6, + min( + time.perf_counter(), + self.measure_end, + ) + - self.measure_start + ) + + prepare_ms = [ + float( + r[ + "prepare_ms" + ] + ) + for r in self.samples + ] + + infer_ms = [ + float( + r[ + "inference_ms" + ] + ) + for r in self.samples + ] + + age_prepare = [ + float( + r[ + "age_at_prepare_start_ms" + ] + ) + for r in self.samples + ] + + wait_after_prepare = [ + float( + r[ + "wait_after_prepare_ms" + ] + ) + for r in self.samples + ] + + age_before_infer = [ + float( + r[ + "age_before_infer_ms" + ] + ) + for r in self.samples + ] + + total_latency = [ + float( + r[ + "capture_to_result_ms" + ] + ) + for r in self.samples + ] + + dropped_before_prepare = int( + self.dropped_before_prepare_measure + ) + + dropped_before_infer = int( + self.dropped_before_infer_measure + ) + + summary = { + "warmup_seconds": float( + self.warmup_s + ), + "measurement_seconds": float( + actual_measure_duration + ), + + "capture_frames_measure": int( + len( + self.capture_times_measure + ) + ), + "prepare_frames_measure": int( + len( + self.prepare_times_measure + ) + ), + "infer_frames_measure": int( + len( + self.samples + ) + ), + + "capture_fps": float( + len( + self.capture_times_measure + ) + / actual_measure_duration + ), + "prepare_fps": float( + len( + self.prepare_times_measure + ) + / actual_measure_duration + ), + "output_fps": float( + len( + self.samples + ) + / actual_measure_duration + ), + + "dropped_before_prepare": dropped_before_prepare, + "dropped_before_infer": dropped_before_infer, + + "prepare_ms": stats_ms( + prepare_ms + ), + "age_at_prepare_start_ms": stats_ms( + age_prepare + ), + "wait_after_prepare_ms": stats_ms( + wait_after_prepare + ), + "age_before_infer_ms": stats_ms( + age_before_infer + ), + "inference_ms": stats_ms( + infer_ms + ), + "capture_to_result_ms": stats_ms( + total_latency + ), + + "errors": list( + self.errors + ), + } + + return summary + + +# ============================================================================= +# Main +# ============================================================================= + +def main(): + ap = argparse.ArgumentParser( + description=( + "Benchmark de campo do modelo frontal: " + "PyTorch, ONNX CUDA, TensorRT e OAK-D threaded." + ) + ) + + ap.add_argument( + "--config", + default="config.json", + ) + + ap.add_argument( + "--test_script", + default=None, + help="Default: test_corridor_agri_v2.py ao lado.", + ) + + ap.add_argument( + "--export_script", + default=None, + help="Default: export_corridor_agri_onnx_v2.py ao lado.", + ) + + ap.add_argument( + "--ckpt", + default=None, + ) + + ap.add_argument( + "--checkpoint", + default="operational", + choices=[ + "operational", + "nav", + "status", + "last", + ], + ) + + ap.add_argument( + "--onnx", + default=None, + help=( + "Default: .field_v2.onnx" + ), + ) + + ap.add_argument( + "--modes", + default="all", + help=( + "CSV: host_prepare,pytorch_core,pytorch_field," + "onnx_cpu,onnx_cuda,onnx_tensorrt ou all." + ), + ) + + ap.add_argument( + "--iterations", + type=int, + default=200, + ) + + ap.add_argument( + "--warmup", + type=int, + default=30, + ) + + ap.add_argument( + "--source_width", + type=int, + default=1920, + ) + + ap.add_argument( + "--source_height", + type=int, + default=1080, + ) + + ap.add_argument( + "--device", + default="cuda", + choices=[ + "cuda", + "cpu", + ], + ) + + ap.add_argument( + "--trt_fp16", + action="store_true", + default=True, + ) + + ap.add_argument( + "--trt_fp32", + action="store_true", + help="Desliga FP16 do TensorRT.", + ) + + ap.add_argument( + "--trt_cache", + default=None, + ) + + ap.add_argument( + "--clear_trt_cache", + action="store_true", + help="Mede cold-start removendo o cache TRT antes de criar sessão.", + ) + + ap.add_argument( + "--camera", + action="store_true", + help="Adiciona benchmark threaded real com OAK-D Lite.", + ) + + ap.add_argument( + "--pipeline_provider", + default="tensorrt", + choices=[ + "cpu", + "cuda", + "tensorrt", + ], + ) + + ap.add_argument( + "--camera_fps", + type=float, + default=30.0, + ) + + ap.add_argument( + "--camera_res", + default="1080p", + choices=[ + "1080p", + "720p", + ], + ) + + ap.add_argument( + "--camera_warmup", + type=float, + default=5.0, + ) + + ap.add_argument( + "--camera_duration", + type=float, + default=30.0, + ) + + ap.add_argument( + "--out_root", + default="benchmark_corridor", + ) + + ap.add_argument( + "--run_name", + default=None, + ) + + args = ap.parse_args() + + if args.iterations <= 0: + ap.error( + "--iterations deve ser > 0." + ) + + if args.warmup < 0: + ap.error( + "--warmup deve ser >= 0." + ) + + modes = parse_modes( + args.modes + ) + + here = Path( + __file__ + ).resolve() + + config_path = Path( + args.config + ).resolve() + + if not config_path.is_file(): + raise FileNotFoundError( + config_path + ) + + config = safe_json_load( + config_path + ) + + project_root = resolve_project_root( + config_path + ) + + test_script = ( + Path( + args.test_script + ).resolve() + if args.test_script + else here.with_name( + "test_corridor_agri_v2.py" + ) + ) + + export_script = ( + Path( + args.export_script + ).resolve() + if args.export_script + else here.with_name( + "export_corridor_agri_onnx_v2.py" + ) + ) + + test_mod = load_module_from_file( + test_script, + "agrobot_test_corridor_v2_bench", + ) + + export_mod = load_module_from_file( + export_script, + "agrobot_export_corridor_v2_bench", + ) + + if args.ckpt: + checkpoint_path = Path( + args.ckpt + ).resolve() + else: + checkpoint_path = test_mod.default_checkpoint_path( + project_root, + config, + args.checkpoint, + ) + + if args.onnx: + onnx_path = Path( + args.onnx + ).resolve() + else: + onnx_path = ( + checkpoint_path.parent + / f"{checkpoint_path.stem}.field_v2.onnx" + ) + + run_name = ( + args.run_name + or ( + f"{checkpoint_path.stem}_" + f"{timestamp_run_name()}" + ) + ) + + out_root = Path( + args.out_root + ) + + if not out_root.is_absolute(): + out_root = ( + project_root + / out_root + ) + + run_dir = ( + out_root + / run_name + ).resolve() + + run_dir.mkdir( + parents=True, + exist_ok=True, + ) + + environment = collect_environment() + + save_json( + run_dir + / "environment.json", + environment, + ) + + use_cuda = ( + args.device + == "cuda" + and torch.cuda.is_available() + ) + + if ( + args.device + == "cuda" + and not use_cuda + ): + print( + "[WARN] CUDA solicitado mas indisponível. " + "PyTorch benchmark usará CPU." + ) + + device = torch.device( + "cuda" + if use_cuda + else "cpu" + ) + + # ------------------------------------------------------------------------- + # Load checkpoint only if required. + # ------------------------------------------------------------------------- + + needs_pytorch = any( + mode.startswith( + "pytorch_" + ) + for mode in modes + ) + + base_model = None + status_head = None + contract = None + checkpoint_raw = None + field_wrapper = None + + # Mesmo quando só ONNX/camera, contract vem do checkpoint para geometria. + ( + base_model_loaded, + status_head_loaded, + contract, + checkpoint_raw, + ) = test_mod.build_runtime_from_checkpoint( + checkpoint_path, + config, + device, + ) + + base_model = base_model_loaded + status_head = status_head_loaded + + W, H = contract.resolution_wh + + # Wrapper exatamente igual ao export. + field_wrapper = export_mod.CorridorFieldOnnx( + base_model=base_model, + status_head=status_head, + resolution_wh=contract.resolution_wh, + norm_mean=contract.norm_mean.tolist() + if hasattr( + contract.norm_mean, + "tolist", + ) + else contract.norm_mean, + norm_std=contract.norm_std.tolist() + if hasattr( + contract.norm_std, + "tolist", + ) + else contract.norm_std, + ).to( + device + ).eval() + + print("=" * 100) + print( + "AGROBOT | CORRIDOR FIELD BENCHMARK V2" + ) + print("=" * 100) + print( + f"Run : {run_name}" + ) + print( + f"Checkpoint : {checkpoint_path}" + ) + print( + f"ONNX : {onnx_path}" + ) + print( + f"Resolution : {W}x{H}" + ) + print( + f"Source synth : " + f"{args.source_width}x{args.source_height}" + ) + print( + f"Device PyTorch : {device}" + ) + print( + f"Modes : {modes}" + ) + print( + f"Iterations/warmup : " + f"{args.iterations}/{args.warmup}" + ) + print( + f"Camera : {args.camera}" + ) + print("=" * 100) + + if any( + mode.startswith( + "onnx_" + ) + for mode in modes + ) or args.camera: + if not onnx_path.is_file(): + raise FileNotFoundError( + f"ONNX não encontrado: {onnx_path}\n" + "Rode export_corridor_agri_onnx_v2.py primeiro." + ) + + # ------------------------------------------------------------------------- + # Build reference input. + # ------------------------------------------------------------------------- + + ref_image_path = find_reference_image( + project_root, + contract.resolution_wh, + ) + + if ref_image_path is not None: + ref_bgr = cv2.imread( + str( + ref_image_path + ), + cv2.IMREAD_COLOR, + ) + + if ref_bgr is None: + ref_bgr = make_synthetic_source_frame( + ( + args.source_width, + args.source_height, + ) + ) + + ref_source = "synthetic_raw" + else: + ref_source = str( + ref_image_path + ) + else: + ref_bgr = make_synthetic_source_frame( + ( + args.source_width, + args.source_height, + ) + ) + + ref_source = "synthetic_raw" + + # Para host_prepare precisamos sempre uma source grande. + synthetic_raw = make_synthetic_source_frame( + ( + args.source_width, + args.source_height, + ) + ) + + input_np = prepare_rgb_u8_nhwc( + ref_bgr, + resolution_wh=contract.resolution_wh, + roi_inicio=contract.roi_inicio, + roi_tamanho=contract.roi_tamanho, + ) + + print( + f"[INPUT] referência={ref_source}" + ) + print( + f"[INPUT] contract={input_np.shape} {input_np.dtype}" + ) + + results: Dict[ + str, + dict + ] = {} + + csv_rows = [] + + # ------------------------------------------------------------------------- + # Host prepare + # ------------------------------------------------------------------------- + + if "host_prepare" in modes: + def host_prepare_call(): + return prepare_rgb_u8_nhwc( + synthetic_raw, + resolution_wh=contract.resolution_wh, + roi_inicio=contract.roi_inicio, + roi_tamanho=contract.roi_tamanho, + ) + + result = benchmark_callable( + name="host_prepare", + call=host_prepare_call, + warmup=args.warmup, + iterations=args.iterations, + ) + + result[ + "source_shape" + ] = list( + synthetic_raw.shape + ) + + result[ + "target_shape" + ] = list( + input_np.shape + ) + + results[ + "host_prepare" + ] = result + + csv_rows.append( + flatten_stats_row( + "host_prepare", + "host", + result[ + "latency_ms" + ], + { + "first_ms": result[ + "first_ms" + ], + }, + ) + ) + + # ------------------------------------------------------------------------- + # PyTorch core + # ------------------------------------------------------------------------- + + x_norm = uint8_contract_to_normalized_nchw( + input_np, + mean=( + contract.norm_mean.tolist() + if hasattr( + contract.norm_mean, + "tolist", + ) + else contract.norm_mean + ), + std=( + contract.norm_std.tolist() + if hasattr( + contract.norm_std, + "tolist", + ) + else contract.norm_std + ), + device=device, + ) + + if "pytorch_core" in modes: + def pt_core_call(): + return pytorch_core_call( + base_model, + status_head, + x_norm, + contract.resolution_wh, + ) + + result = benchmark_callable( + name="pytorch_core", + call=pt_core_call, + warmup=args.warmup, + iterations=args.iterations, + sync_before_after=( + lambda: synchronize_cuda( + device + ) + ), + ) + + result[ + "input_shape" + ] = list( + x_norm.shape + ) + + result[ + "input_dtype" + ] = str( + x_norm.dtype + ) + + results[ + "pytorch_core" + ] = result + + csv_rows.append( + flatten_stats_row( + "pytorch_core", + "model", + result[ + "latency_ms" + ], + { + "first_ms": result[ + "first_ms" + ], + }, + ) + ) + + # ------------------------------------------------------------------------- + # PyTorch field wrapper + # ------------------------------------------------------------------------- + + if "pytorch_field" in modes: + x_u8 = torch.from_numpy( + np.ascontiguousarray( + input_np + ) + ).to( + device, + non_blocking=True, + ) + + def pt_field_call(): + with torch.inference_mode(): + return field_wrapper( + x_u8 + ) + + result = benchmark_callable( + name="pytorch_field", + call=pt_field_call, + warmup=args.warmup, + iterations=args.iterations, + sync_before_after=( + lambda: synchronize_cuda( + device + ) + ), + ) + + result[ + "input_shape" + ] = list( + x_u8.shape + ) + + result[ + "input_dtype" + ] = str( + x_u8.dtype + ) + + results[ + "pytorch_field" + ] = result + + csv_rows.append( + flatten_stats_row( + "pytorch_field", + "field_reference", + result[ + "latency_ms" + ], + { + "first_ms": result[ + "first_ms" + ], + }, + ) + ) + + # ------------------------------------------------------------------------- + # Release PyTorch before ORT/TensorRT. + # + # Não queremos medir TensorRT competindo por VRAM com uma segunda cópia + # inteira do mesmo modelo carregada no PyTorch. + # ------------------------------------------------------------------------- + + if any( + mode.startswith( + "onnx_" + ) + for mode in modes + ) or args.camera: + try: + del x_norm + except Exception: + pass + + try: + del x_u8 + except Exception: + pass + + try: + del field_wrapper + except Exception: + pass + + try: + del base_model + except Exception: + pass + + try: + del status_head + except Exception: + pass + + gc.collect() + + if torch.cuda.is_available(): + torch.cuda.empty_cache() + torch.cuda.synchronize() + + print( + "[MEM] PyTorch runtime liberado antes dos benchmarks ORT/TensorRT." + ) + + # ------------------------------------------------------------------------- + # TRT cache + # ------------------------------------------------------------------------- + + trt_cache_dir = ( + Path( + args.trt_cache + ).resolve() + if args.trt_cache + else ( + onnx_path.parent + / "trt_cache_field_v2" + ) + ) + + if args.clear_trt_cache: + print( + f"[TRT] removendo cache para cold-start: {trt_cache_dir}" + ) + + if trt_cache_dir.exists(): + shutil.rmtree( + trt_cache_dir + ) + + trt_fp16 = ( + not bool( + args.trt_fp32 + ) + ) + + ort_sessions = {} + + # ------------------------------------------------------------------------- + # ONNX providers + # ------------------------------------------------------------------------- + + provider_modes = [] + + if "onnx_cpu" in modes: + provider_modes.append( + "cpu" + ) + + if "onnx_cuda" in modes: + provider_modes.append( + "cuda" + ) + + if "onnx_tensorrt" in modes: + provider_modes.append( + "tensorrt" + ) + + for provider_mode in provider_modes: + try: + result, session = benchmark_ort( + onnx_path=onnx_path, + provider_mode=provider_mode, + input_np=input_np, + warmup=args.warmup, + iterations=args.iterations, + trt_cache_dir=trt_cache_dir, + trt_fp16=trt_fp16, + ) + + key = ( + f"onnx_{provider_mode}" + ) + + results[ + key + ] = result + + ort_sessions[ + provider_mode + ] = session + + csv_rows.append( + flatten_stats_row( + key, + "onnxruntime", + result[ + "latency_ms" + ], + { + "first_ms": result[ + "first_ms" + ], + "session_create_seconds": result[ + "provider" + ][ + "session_create_seconds" + ], + "active_providers": "|".join( + result[ + "provider" + ][ + "active_providers" + ] + ), + }, + ) + ) + + except Exception as exc: + key = ( + f"onnx_{provider_mode}" + ) + + print( + f"[WARN] {key} indisponível/falhou: " + f"{type(exc).__name__}: {exc}" + ) + + results[ + key + ] = { + "error": ( + f"{type(exc).__name__}: {exc}" + ) + } + + csv_rows.append({ + "name": key, + "kind": "onnxruntime", + "error": results[ + key + ][ + "error" + ], + }) + + # ------------------------------------------------------------------------- + # Camera threaded pipeline + # ------------------------------------------------------------------------- + + camera_summary = None + camera_samples = [] + + if args.camera: + provider_mode = str( + args.pipeline_provider + ) + + print() + print("=" * 100) + print( + f"CAMERA PIPELINE | provider={provider_mode}" + ) + print("=" * 100) + + session = ort_sessions.get( + provider_mode + ) + + session_info = None + + if session is None: + session, session_info = create_ort_session( + onnx_path=onnx_path, + provider_mode=provider_mode, + trt_cache_dir=trt_cache_dir, + trt_fp16=trt_fp16, + ) + else: + session_info = results.get( + f"onnx_{provider_mode}", + {}, + ).get( + "provider" + ) + + cam_bench = CameraPipelineBenchmark( + session=session, + contract=contract, + camera_fps=args.camera_fps, + camera_resolution=args.camera_res, + warmup_s=args.camera_warmup, + duration_s=args.camera_duration, + ) + + camera_summary = cam_bench.run() + + camera_summary[ + "provider" + ] = session_info + + camera_summary[ + "camera_resolution" + ] = args.camera_res + + camera_summary[ + "camera_requested_fps" + ] = float( + args.camera_fps + ) + + camera_samples = cam_bench.samples + + results[ + "camera_pipeline" + ] = camera_summary + + print() + print( + "[CAM RESULT] " + f"capture={camera_summary['capture_fps']:.2f} FPS | " + f"prepare={camera_summary['prepare_fps']:.2f} FPS | " + f"output={camera_summary['output_fps']:.2f} FPS" + ) + + if camera_summary[ + "inference_ms" + ][ + "mean" + ] is not None: + print( + "[CAM RESULT] " + f"infer mean/p95=" + f"{camera_summary['inference_ms']['mean']:.2f}/" + f"{camera_summary['inference_ms']['p95']:.2f} ms" + ) + + if camera_summary[ + "capture_to_result_ms" + ][ + "mean" + ] is not None: + print( + "[CAM RESULT] " + f"E2E capture→result mean/p95=" + f"{camera_summary['capture_to_result_ms']['mean']:.2f}/" + f"{camera_summary['capture_to_result_ms']['p95']:.2f} ms" + ) + + print( + "[CAM RESULT] " + f"drops raw/prepared=" + f"{camera_summary['dropped_before_prepare']}/" + f"{camera_summary['dropped_before_infer']}" + ) + + if camera_summary[ + "errors" + ]: + print( + f"[CAM WARN] errors={camera_summary['errors']}" + ) + + csv_rows.append({ + "name": "camera_pipeline", + "kind": "threaded_camera", + "capture_fps": camera_summary[ + "capture_fps" + ], + "prepare_fps": camera_summary[ + "prepare_fps" + ], + "output_fps": camera_summary[ + "output_fps" + ], + "prepare_mean_ms": camera_summary[ + "prepare_ms" + ][ + "mean" + ], + "prepare_p95_ms": camera_summary[ + "prepare_ms" + ][ + "p95" + ], + "infer_mean_ms": camera_summary[ + "inference_ms" + ][ + "mean" + ], + "infer_p95_ms": camera_summary[ + "inference_ms" + ][ + "p95" + ], + "e2e_mean_ms": camera_summary[ + "capture_to_result_ms" + ][ + "mean" + ], + "e2e_p95_ms": camera_summary[ + "capture_to_result_ms" + ][ + "p95" + ], + "dropped_before_prepare": camera_summary[ + "dropped_before_prepare" + ], + "dropped_before_infer": camera_summary[ + "dropped_before_infer" + ], + }) + + # ------------------------------------------------------------------------- + # GPU memory snapshot + # ------------------------------------------------------------------------- + + gpu_memory = None + + if torch.cuda.is_available(): + synchronize_cuda( + torch.device( + "cuda" + ) + ) + + gpu_memory = { + "allocated_bytes": int( + torch.cuda.memory_allocated( + 0 + ) + ), + "reserved_bytes": int( + torch.cuda.memory_reserved( + 0 + ) + ), + "max_allocated_bytes": int( + torch.cuda.max_memory_allocated( + 0 + ) + ), + "max_reserved_bytes": int( + torch.cuda.max_memory_reserved( + 0 + ) + ), + } + + # ------------------------------------------------------------------------- + # Summary + # ------------------------------------------------------------------------- + + summary = { + "schema": "agrobot.visual.corridor.benchmark.v2", + "benchmark_version": BENCHMARK_VERSION, + "created_at": now_iso(), + + "run_name": run_name, + + "project_root": str( + project_root + ), + "config": str( + config_path + ), + "checkpoint": str( + checkpoint_path + ), + "onnx": str( + onnx_path + ), + + "contract": { + "resolution_wh": [ + int( + W + ), + int( + H + ), + ], + "roi_inicio": float( + contract.roi_inicio + ), + "roi_tamanho": float( + contract.roi_tamanho + ), + "nav_class_id": int( + contract.nav_class_id + ), + "non_nav_class_id": int( + contract.non_nav_class_id + ), + "status_classes": { + str( + k + ): v + for k, v in contract.label_name_by_id.items() + }, + "input_name": INPUT_NAME, + "outputs": [ + SEG_OUTPUT_NAME, + STATUS_OUTPUT_NAME, + ], + }, + + "settings": { + "modes": modes, + "iterations": int( + args.iterations + ), + "warmup": int( + args.warmup + ), + "trt_fp16": bool( + trt_fp16 + ), + "trt_cache_dir": str( + trt_cache_dir + ), + "clear_trt_cache": bool( + args.clear_trt_cache + ), + "camera": bool( + args.camera + ), + "pipeline_provider": str( + args.pipeline_provider + ), + }, + + "reference_input": { + "source": ref_source, + "shape": list( + input_np.shape + ), + "dtype": str( + input_np.dtype + ), + }, + + "environment": environment, + "gpu_memory": gpu_memory, + + "results": results, + } + + save_json( + run_dir + / "benchmark_summary.json", + summary, + ) + + write_csv( + run_dir + / "benchmark_results.csv", + csv_rows, + ) + + if camera_samples: + write_csv( + run_dir + / "camera_pipeline_samples.csv", + camera_samples, + ) + + # ------------------------------------------------------------------------- + # Compact ranking + # ------------------------------------------------------------------------- + + print() + print("=" * 100) + print( + "BENCHMARK FINALIZADO" + ) + print("=" * 100) + + ranking = [] + + for name, result in results.items(): + if ( + isinstance( + result, + dict + ) + and "latency_ms" + in result + and result[ + "latency_ms" + ].get( + "mean" + ) + is not None + ): + ranking.append( + ( + name, + float( + result[ + "latency_ms" + ][ + "mean" + ] + ), + ) + ) + + for name, mean_ms in sorted( + ranking, + key=lambda x: x[ + 1 + ], + ): + print( + f"{name:22s} " + f"{mean_ms:9.3f} ms " + f"{1000.0 / mean_ms:8.2f} FPS" + ) + + if camera_summary is not None: + print("-" * 100) + print( + f"{'camera_pipeline':22s} " + f"output={camera_summary['output_fps']:.2f} FPS | " + f"E2E p50={camera_summary['capture_to_result_ms']['p50']:.2f} ms | " + f"p95={camera_summary['capture_to_result_ms']['p95']:.2f} ms" + ) + + print("-" * 100) + print( + f"Summary : {run_dir / 'benchmark_summary.json'}" + ) + print( + f"CSV : {run_dir / 'benchmark_results.csv'}" + ) + + if camera_samples: + print( + f"Camera : {run_dir / 'camera_pipeline_samples.csv'}" + ) + + print("=" * 100) + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/oak-d/audit/review_corridor_dataset_v2.py b/Python/OAK/datasets/oak-d/audit/review_corridor_dataset_v2.py new file mode 100644 index 000000000..38e876314 --- /dev/null +++ b/Python/OAK/datasets/oak-d/audit/review_corridor_dataset_v2.py @@ -0,0 +1,5125 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +review_corridor_dataset_v2.py +============================= + +Auditoria / model-mining do dataset frontal de corredor. + +Objetivo +-------- +Rodar o checkpoint treinado sobre TODO o dataset normalizado, comparar: + + 1) segmentação GT vs predição; + 2) estado global GT vs segunda cabeça; + +e ranquear as amostras que mais merecem revisão humana. + +IMPORTANTE +---------- +Discordância modelo x GT NÃO prova que a anotação está errada. + +O modelo é usado como um "segundo revisor": + - conflito de alta confiança no interior de uma região; + - falso navegável de alta confiança; + - estado global diferente com alta confiança; + - composição semântica incompatível com o grupo anotado; + +são sinais úteis para encontrar: + - máscaras antigas ruins; + - labels globais incorretos; + - inconsistências de anotação; + - ou, igualmente importante, falhas sistemáticas do próprio modelo. + +O review_score é SOMENTE um ranking heurístico. +Não é probabilidade estatística de erro humano. + +Backend +------- +Este auditor usa o checkpoint PyTorch v2 de propósito. + +O ONNX de CAMPO v2 foi desenhado para ser mínimo: + rgb_u8_nhwc -> seg_ids + label_probs + +Ele não expõe semantic logits/probabilities por pixel, portanto não é adequado +para "high-confidence model mining". + +Não vamos poluir o contrato de campo só para auditoria. Se futuramente quisermos +auditoria ONNX acelerada, criaremos um export AUDIT separado. + +Dependência +----------- +Por padrão carrega dinamicamente: + test_corridor_agri_v2.py + +Assim reutilizamos EXATAMENTE: + - reconstrução SegFormer; + - CorridorStatusHead 3x4; + - contrato do checkpoint; + - preprocess; + - labelmap; + - descoberta de dataset. + +Entrada padrão +-------------- + dataset/x/group/ + +Estrutura esperada: + group/ + navegavel/ + images/ + masks/ + labels/ + naonavegavel/ + ... + naonavegave_navegavel/ + ... + +Também aceita split/train ou split/val. + +Saídas +------ + dataset/revisao_corridor/ + reports// + review_report.csv + review_candidates.csv + review_summary.json + confusion_semantic.csv + confusion_status.csv + README_REVIEW.txt + + group// + original_images/ + original_masks/ + original_labels/ + predictions/ + panels/ # opcional + final_masks/ # PRESERVADO pelo --clear_review + review_order.csv + + review_source_manifest.csv + +Exemplos +-------- +1) Auditoria completa + export dos casos tecnicamente suspeitos: + + python review_corridor_dataset_v2.py ^ + --clear_review ^ + --save_panels + +2) Só gerar relatórios: + + python review_corridor_dataset_v2.py ^ + --report_only ^ + --run_name best_operational_audit + +3) Pegar os 25 casos mais suspeitos de cada grupo x estado: + + python review_corridor_dataset_v2.py ^ + --top_k_per_stratum 25 ^ + --clear_review ^ + --save_panels + +4) Exportar somente score >= 70/100: + + python review_corridor_dataset_v2.py ^ + --export_mode score ^ + --min_suspicion_pct 70 ^ + --clear_review + +Critérios principais +-------------------- +SEGMENTAÇÃO: + semantic_disagree_frac + highconf_disagree_frac + interior_disagree_frac + highconf_interior_disagree_frac + unsafe_nav_frac + highconf_unsafe_nav_frac + blocked_nav_frac + nav_iou + non_nav_iou + +SEGUNDA CABEÇA: + gt_status + pred_status + status_confidence + gt_status_probability + status_margin + status_mismatch + +COMPOSIÇÃO: + gt_group + pred_group + group_mismatch + +Fronteira vs interior +--------------------- +Discordâncias muito perto da borda da máscara recebem menos peso no ranking. +Conflitos de alta confiança LONGE da borda recebem mais peso. + +Isso é importante para separar: + "1 ou 2 px de contorno" +de: + "o modelo vê uma região inteira diferente da anotação". +""" + +from __future__ import annotations + +import argparse +import csv +import gc +import importlib.util +import json +import math +import os +import re +import shutil +import sys +import time +from collections import Counter, defaultdict +from dataclasses import dataclass +from pathlib import Path +from typing import Dict, Iterable, List, Optional, Sequence, Tuple + +import cv2 +import numpy as np +import torch +import torch.nn.functional as F + + +# ============================================================================= +# Constants +# ============================================================================= + +REVIEWER_VERSION = "agrobot_corridor_dataset_review_v2.0" +REVIEW_SCHEMA = "agrobot.visual.corridor.dataset_review.v2" +IGNORE_INDEX = 255 + +IMAGE_EXTS = { + ".png", + ".jpg", + ".jpeg", + ".bmp", + ".webp", + ".tif", + ".tiff", +} + + +# ============================================================================= +# Generic helpers +# ============================================================================= + +def now_iso() -> str: + from datetime import datetime + + return datetime.now().isoformat( + timespec="seconds" + ) + + +def load_json( + path: Path, +) -> dict: + with path.open( + "r", + encoding="utf-8", + ) as f: + data = json.load( + f + ) + + if not isinstance( + data, + dict, + ): + raise ValueError( + f"Esperava objeto JSON: {path}" + ) + + return data + + +def save_json( + path: Path, + data: dict, +) -> None: + path.parent.mkdir( + parents=True, + exist_ok=True, + ) + + tmp = path.with_suffix( + path.suffix + + ".tmp" + ) + + with tmp.open( + "w", + encoding="utf-8", + ) as f: + json.dump( + data, + f, + ensure_ascii=False, + indent=2, + ) + + os.replace( + tmp, + path, + ) + + +def write_csv( + path: Path, + rows: Sequence[dict], + fieldnames: Optional[ + Sequence[str] + ] = None, +) -> None: + path.parent.mkdir( + parents=True, + exist_ok=True, + ) + + if fieldnames is None: + keys = [] + seen = set() + + for row in rows: + for key in row: + if key not in seen: + seen.add( + key + ) + keys.append( + key + ) + + fieldnames = keys + + with path.open( + "w", + newline="", + encoding="utf-8-sig", + ) as f: + writer = csv.DictWriter( + f, + fieldnames=list( + fieldnames + ), + extrasaction="ignore", + ) + + writer.writeheader() + + for row in rows: + writer.writerow( + row + ) + + +def safe_float( + value, + default: float = 0.0, +) -> float: + try: + x = float( + value + ) + + return ( + x + if math.isfinite( + x + ) + else default + ) + except Exception: + return default + + +def clear_dir( + path: Path, +) -> None: + if path.exists(): + shutil.rmtree( + path + ) + + path.mkdir( + parents=True, + exist_ok=True, + ) + + +def natural_key( + text: str, +): + parts = re.split( + r"(\d+)", + str( + text + ), + ) + + return [ + int( + p + ) + if p.isdigit() + else p.lower() + for p in parts + ] + + +# ============================================================================= +# Load shared tester +# ============================================================================= + +def load_test_module( + script_path: Path, +): + if not script_path.is_file(): + raise FileNotFoundError( + f"Tester base não encontrado: {script_path}\n" + "Coloque review_corridor_dataset_v2.py ao lado de " + "test_corridor_agri_v2.py ou informe --test_script." + ) + + spec = importlib.util.spec_from_file_location( + "agrobot_corridor_test_v2", + str( + script_path + ), + ) + + if ( + spec is None + or spec.loader is None + ): + raise RuntimeError( + f"Não consegui importar tester: {script_path}" + ) + + module = importlib.util.module_from_spec( + spec + ) + + sys.modules[ + spec.name + ] = module + + spec.loader.exec_module( + module + ) + + required = ( + "safe_json_load", + "find_project_root", + "load_labelmap", + "discover_samples", + "default_checkpoint_path", + "build_runtime_from_checkpoint", + "prepare_input", + "load_gt_mask", + "read_gt_label", + "colorize_ids", + ) + + missing = [ + name + for name in required + if not hasattr( + module, + name, + ) + ] + + if missing: + raise RuntimeError( + f"Tester base incompatível; faltando={missing}" + ) + + return module + + +# ============================================================================= +# Metrics +# ============================================================================= + +def confusion_matrix_np( + pred: np.ndarray, + gt: np.ndarray, + num_classes: int, + ignore_id: int, +) -> np.ndarray: + if ( + pred.shape + != gt.shape + ): + pred = cv2.resize( + pred.astype( + np.uint8 + ), + ( + gt.shape[ + 1 + ], + gt.shape[ + 0 + ], + ), + interpolation=cv2.INTER_NEAREST, + ) + + valid = ( + gt + != int( + ignore_id + ) + ) + + pred_v = pred[ + valid + ].astype( + np.int64 + ) + gt_v = gt[ + valid + ].astype( + np.int64 + ) + + valid2 = ( + (gt_v >= 0) + & ( + gt_v + < num_classes + ) + & ( + pred_v + >= 0 + ) + & ( + pred_v + < num_classes + ) + ) + + pred_v = pred_v[ + valid2 + ] + gt_v = gt_v[ + valid2 + ] + + cm = np.zeros( + ( + num_classes, + num_classes, + ), + dtype=np.int64, + ) + + if gt_v.size: + idx = ( + gt_v + * num_classes + + pred_v + ) + + bins = np.bincount( + idx, + minlength=( + num_classes + * num_classes + ), + ) + + cm += bins.reshape( + num_classes, + num_classes, + ) + + return cm + + +def binary_metrics_from_cm( + cm: np.ndarray, + nav_id: int, + non_nav_id: int, +) -> dict: + cmf = cm.astype( + np.float64 + ) + + tp = np.diag( + cmf + ) + fp = cmf.sum( + axis=0 + ) - tp + fn = cmf.sum( + axis=1 + ) - tp + + iou = tp / np.maximum( + tp + + fp + + fn, + 1e-12, + ) + + precision = tp / np.maximum( + tp + + fp, + 1e-12, + ) + + recall = tp / np.maximum( + tp + + fn, + 1e-12, + ) + + f1 = ( + 2.0 + * precision + * recall + / np.maximum( + precision + + recall, + 1e-12, + ) + ) + + acc = float( + tp.sum() + / max( + 1.0, + cmf.sum(), + ) + ) + + gt_non_nav = float( + cmf[ + non_nav_id, + :, + ].sum() + ) + + unsafe = float( + cmf[ + non_nav_id, + nav_id, + ] + / max( + 1.0, + gt_non_nav, + ) + ) + + gt_nav = float( + cmf[ + nav_id, + :, + ].sum() + ) + + blocked = float( + cmf[ + nav_id, + non_nav_id, + ] + / max( + 1.0, + gt_nav, + ) + ) + + return { + "acc": acc, + "miou": float( + np.mean( + iou + ) + ), + "nav_iou": float( + iou[ + nav_id + ] + ), + "non_nav_iou": float( + iou[ + non_nav_id + ] + ), + "nav_f1": float( + f1[ + nav_id + ] + ), + "non_nav_f1": float( + f1[ + non_nav_id + ] + ), + "unsafe_nav_rate": unsafe, + "blocked_nav_rate": blocked, + } + + +def status_metrics_from_cm( + cm: np.ndarray, +) -> dict: + cmf = cm.astype( + np.float64 + ) + + tp = np.diag( + cmf + ) + fp = cmf.sum( + axis=0 + ) - tp + fn = cmf.sum( + axis=1 + ) - tp + support = cmf.sum( + axis=1 + ) + + precision = tp / np.maximum( + tp + + fp, + 1e-12, + ) + recall = tp / np.maximum( + tp + + fn, + 1e-12, + ) + f1 = ( + 2.0 + * precision + * recall + / np.maximum( + precision + + recall, + 1e-12, + ) + ) + + present = support > 0 + + return { + "acc": float( + tp.sum() + / max( + 1.0, + cmf.sum(), + ) + ), + "macro_f1": ( + float( + f1[ + present + ].mean() + ) + if np.any( + present + ) + else 0.0 + ), + "precision": precision.tolist(), + "recall": recall.tolist(), + "f1": f1.tolist(), + "support": support.tolist(), + } + + +def write_cm_csv( + path: Path, + cm: np.ndarray, + labels: Sequence[str], +) -> None: + rows = [] + + for gi, gt_name in enumerate( + labels + ): + row = { + "gt\\pred": gt_name + } + + for pi, pred_name in enumerate( + labels + ): + row[ + pred_name + ] = int( + cm[ + gi, + pi, + ] + ) + + rows.append( + row + ) + + write_csv( + path, + rows, + [ + "gt\\pred", + *labels, + ], + ) + + +# ============================================================================= +# Boundary analysis +# ============================================================================= + +def make_boundary_band( + gt: np.ndarray, + ignore_id: int, + radius_px: int, +) -> np.ndarray: + """ + Banda ao redor das transições de classe do GT. + + A ideia é não tratar erro de 1..N px na borda com o mesmo peso de uma + discordância no interior de uma região. + """ + valid = ( + gt + != int( + ignore_id + ) + ) + + if not np.any( + valid + ): + return np.zeros_like( + gt, + dtype=bool, + ) + + # Gradiente binário de classe. + # Para duas classes, qualquer vizinho diferente marca fronteira. + gt_u8 = gt.astype( + np.uint8, + copy=False, + ) + + boundary = np.zeros_like( + valid, + dtype=np.uint8, + ) + + neighbors = ( + (-1, 0), + (1, 0), + (0, -1), + (0, 1), + (-1, -1), + (-1, 1), + (1, -1), + (1, 1), + ) + + for dy, dx in neighbors: + shifted = np.roll( + gt_u8, + shift=( + dy, + dx, + ), + axis=( + 0, + 1, + ), + ) + + shifted_valid = np.roll( + valid, + shift=( + dy, + dx, + ), + axis=( + 0, + 1, + ), + ) + + diff = ( + valid + & shifted_valid + & ( + shifted + != gt_u8 + ) + ) + + boundary[ + diff + ] = 1 + + if radius_px > 0: + k = ( + 2 + * int( + radius_px + ) + + 1 + ) + + kernel = np.ones( + ( + k, + k, + ), + dtype=np.uint8, + ) + + boundary = cv2.dilate( + boundary, + kernel, + iterations=1, + ) + + return ( + boundary.astype( + bool + ) + & valid + ) + + +# ============================================================================= +# Semantic signals +# ============================================================================= + +def derive_composition_group( + pred: np.ndarray, + nav_id: int, + non_nav_id: int, + min_fraction: float, +) -> Tuple[str, float]: + valid = ( + (pred == nav_id) + | ( + pred + == non_nav_id + ) + ) + + n = int( + valid.sum() + ) + + if n <= 0: + return ( + "unknown", + 0.0, + ) + + nav_frac = float( + ( + pred[ + valid + ] + == nav_id + ).mean() + ) + + if nav_frac <= float( + min_fraction + ): + group = "naonavegavel" + + elif ( + 1.0 + - nav_frac + ) <= float( + min_fraction + ): + group = "navegavel" + + else: + group = "naonavegave_navegavel" + + return ( + group, + nav_frac, + ) + + +def semantic_sample_signals( + *, + pred: np.ndarray, + probs: np.ndarray, + gt: np.ndarray, + nav_id: int, + non_nav_id: int, + ignore_id: int, + high_confidence: float, + uncertainty_margin: float, + boundary_radius: int, +) -> dict: + if pred.shape != gt.shape: + pred = cv2.resize( + pred.astype( + np.uint8 + ), + ( + gt.shape[ + 1 + ], + gt.shape[ + 0 + ], + ), + interpolation=cv2.INTER_NEAREST, + ) + + if probs.shape[ + 1: + ] != gt.shape: + resized = [] + + for c in range( + probs.shape[ + 0 + ] + ): + resized.append( + cv2.resize( + probs[ + c + ].astype( + np.float32 + ), + ( + gt.shape[ + 1 + ], + gt.shape[ + 0 + ], + ), + interpolation=cv2.INTER_LINEAR, + ) + ) + + probs = np.stack( + resized, + axis=0, + ) + + valid = ( + gt + != int( + ignore_id + ) + ) + + valid_n = int( + valid.sum() + ) + + if valid_n <= 0: + raise RuntimeError( + "Amostra sem pixels GT válidos." + ) + + top1 = np.max( + probs, + axis=0, + ) + + if probs.shape[ + 0 + ] >= 2: + part = np.partition( + probs, + kth=( + probs.shape[ + 0 + ] + - 2 + ), + axis=0, + ) + top2 = part[ + -2 + ] + else: + top2 = np.zeros_like( + top1 + ) + + margin = ( + top1 + - top2 + ) + + wrong = ( + valid + & ( + pred + != gt + ) + ) + + correct = ( + valid + & ( + pred + == gt + ) + ) + + highconf = ( + top1 + >= float( + high_confidence + ) + ) + + highconf_wrong = ( + wrong + & highconf + ) + + uncertain = ( + valid + & ( + margin + <= float( + uncertainty_margin + ) + ) + ) + + boundary = make_boundary_band( + gt, + ignore_id, + boundary_radius, + ) + + interior = ( + valid + & ~boundary + ) + + interior_wrong = ( + wrong + & interior + ) + + hc_interior_wrong = ( + highconf_wrong + & interior + ) + + unsafe = ( + valid + & ( + gt + == non_nav_id + ) + & ( + pred + == nav_id + ) + ) + + blocked = ( + valid + & ( + gt + == nav_id + ) + & ( + pred + == non_nav_id + ) + ) + + hc_unsafe = ( + unsafe + & highconf + ) + + hc_unsafe_interior = ( + hc_unsafe + & interior + ) + + hc_blocked = ( + blocked + & highconf + ) + + hc_blocked_interior = ( + hc_blocked + & interior + ) + + def frac( + mask: np.ndarray, + denom: Optional[ + int + ] = None, + ) -> float: + d = ( + valid_n + if denom is None + else int( + denom + ) + ) + + return float( + int( + mask.sum() + ) + / max( + 1, + d, + ) + ) + + gt_non_nav_n = int( + ( + valid + & ( + gt + == non_nav_id + ) + ).sum() + ) + + gt_nav_n = int( + ( + valid + & ( + gt + == nav_id + ) + ).sum() + ) + + correct_n = int( + correct.sum() + ) + + wrong_n = int( + wrong.sum() + ) + + mean_conf_correct = ( + float( + top1[ + correct + ].mean() + ) + if correct_n + else 0.0 + ) + + mean_conf_wrong = ( + float( + top1[ + wrong + ].mean() + ) + if wrong_n + else 0.0 + ) + + return { + "valid_pixels": valid_n, + + "wrong_pixels": wrong_n, + "semantic_disagree_frac": frac( + wrong + ), + + "boundary_pixels": int( + boundary.sum() + ), + "boundary_frac": frac( + boundary + ), + + "interior_pixels": int( + interior.sum() + ), + "interior_wrong_pixels": int( + interior_wrong.sum() + ), + "interior_disagree_frac": frac( + interior_wrong + ), + + "highconf_wrong_pixels": int( + highconf_wrong.sum() + ), + "highconf_disagree_frac": frac( + highconf_wrong + ), + + "highconf_interior_wrong_pixels": int( + hc_interior_wrong.sum() + ), + "highconf_interior_disagree_frac": frac( + hc_interior_wrong + ), + + "uncertain_pixels": int( + uncertain.sum() + ), + "uncertain_frac": frac( + uncertain + ), + + "mean_conf_correct": mean_conf_correct, + "mean_conf_wrong": mean_conf_wrong, + + "unsafe_nav_pixels": int( + unsafe.sum() + ), + "unsafe_nav_frac_all": frac( + unsafe + ), + "unsafe_nav_frac_gt_nonnav": frac( + unsafe, + gt_non_nav_n, + ), + + "highconf_unsafe_nav_pixels": int( + hc_unsafe.sum() + ), + "highconf_unsafe_nav_frac_all": frac( + hc_unsafe + ), + "highconf_unsafe_nav_frac_gt_nonnav": frac( + hc_unsafe, + gt_non_nav_n, + ), + + "highconf_unsafe_nav_interior_pixels": int( + hc_unsafe_interior.sum() + ), + "highconf_unsafe_nav_interior_frac_all": frac( + hc_unsafe_interior + ), + + "blocked_nav_pixels": int( + blocked.sum() + ), + "blocked_nav_frac_all": frac( + blocked + ), + "blocked_nav_frac_gt_nav": frac( + blocked, + gt_nav_n, + ), + + "highconf_blocked_nav_pixels": int( + hc_blocked.sum() + ), + "highconf_blocked_nav_frac_all": frac( + hc_blocked + ), + + "highconf_blocked_nav_interior_pixels": int( + hc_blocked_interior.sum() + ), + "highconf_blocked_nav_interior_frac_all": frac( + hc_blocked_interior + ), + } + + +# ============================================================================= +# Status signals +# ============================================================================= + +def status_sample_signals( + probs: np.ndarray, + gt_label_id: Optional[int], +) -> dict: + probs = np.asarray( + probs, + dtype=np.float64, + ).reshape( + -1 + ) + + pred_id = int( + np.argmax( + probs + ) + ) + + pred_conf = float( + probs[ + pred_id + ] + ) + + if probs.size >= 2: + sorted_probs = np.sort( + probs + ) + top2 = float( + sorted_probs[ + -2 + ] + ) + else: + top2 = 0.0 + + margin = float( + pred_conf + - top2 + ) + + if ( + gt_label_id is None + or not ( + 0 + <= int( + gt_label_id + ) + < probs.size + ) + ): + gt_prob = None + mismatch = None + else: + gt_prob = float( + probs[ + int( + gt_label_id + ) + ] + ) + mismatch = ( + pred_id + != int( + gt_label_id + ) + ) + + return { + "pred_status_id": pred_id, + "status_confidence": pred_conf, + "status_margin": margin, + "gt_status_probability": gt_prob, + "status_mismatch": ( + int( + mismatch + ) + if mismatch is not None + else None + ), + } + + +# ============================================================================= +# Review score +# ============================================================================= + +def _sat( + value: float, + reference: float, +) -> float: + return float( + min( + max( + float( + value + ) + / max( + float( + reference + ), + 1e-12, + ), + 0.0, + ), + 1.0, + ) + ) + + +def compute_review_score( + row: dict, +) -> Tuple[ + float, + dict, +]: + """ + Score heurístico para RANKING. + + Filosofia: + - alta confiança pesa muito; + - conflito no interior pesa mais que borda; + - falso navegável tem peso operacional maior; + - status global é um eixo independente; + - composição grupo ajuda pouco, apenas como evidência adicional. + + Não é probabilidade. + """ + seg_disagree = _sat( + row.get( + "semantic_disagree_frac", + 0.0, + ), + 0.18, + ) + + hc_disagree = _sat( + row.get( + "highconf_disagree_frac", + 0.0, + ), + 0.035, + ) + + hc_interior = _sat( + row.get( + "highconf_interior_disagree_frac", + 0.0, + ), + 0.018, + ) + + hc_unsafe = _sat( + row.get( + "highconf_unsafe_nav_interior_frac_all", + 0.0, + ), + 0.006, + ) + + seg_score = ( + 0.15 + * seg_disagree + + 0.28 + * hc_disagree + + 0.32 + * hc_interior + + 0.25 + * hc_unsafe + ) + + status_mismatch = bool( + row.get( + "status_mismatch", + 0, + ) + == 1 + ) + + status_conf = safe_float( + row.get( + "status_confidence" + ), + 0.0, + ) + + gt_status_prob = row.get( + "gt_status_probability" + ) + + if status_mismatch: + status_score = min( + 1.0, + ( + 0.40 + + 0.60 + * _sat( + status_conf + - 0.50, + 0.45, + ) + ), + ) + + if ( + gt_status_prob + is not None + ): + status_score = max( + status_score, + _sat( + 0.35 + - float( + gt_status_prob + ), + 0.35, + ), + ) + else: + status_score = 0.0 + + group_score = ( + 1.0 + if int( + row.get( + "group_mismatch", + 0, + ) + ) + == 1 + else 0.0 + ) + + # União suave de evidências. + combined = ( + 1.0 + - ( + 1.0 + - 0.82 + * seg_score + ) + * ( + 1.0 + - 0.68 + * status_score + ) + * ( + 1.0 + - 0.12 + * group_score + ) + ) + + combined = float( + min( + max( + combined, + 0.0, + ), + 1.0, + ) + ) + + parts = { + "seg_review_score": float( + seg_score + ), + "status_review_score": float( + status_score + ), + "group_review_score": float( + group_score + ), + } + + return ( + combined, + parts, + ) + + +def build_reasons( + row: dict, + *, + min_disagree_frac: float, + min_highconf_disagree_frac: float, + min_highconf_interior_frac: float, + min_highconf_unsafe_frac: float, + min_conflict_pixels: int, + status_confidence: float, + min_status_confidence: float, + min_group_disagree_frac: float, +) -> List[ + str +]: + reasons: List[ + str + ] = [] + + if safe_float( + row.get( + "semantic_disagree_frac" + ) + ) >= float( + min_disagree_frac + ): + reasons.append( + "SEMANTIC_DISAGREE" + ) + + hc_pixels = int( + row.get( + "highconf_wrong_pixels", + 0, + ) + ) + + if ( + safe_float( + row.get( + "highconf_disagree_frac" + ) + ) + >= float( + min_highconf_disagree_frac + ) + and hc_pixels + >= int( + min_conflict_pixels + ) + ): + reasons.append( + "HIGH_CONF_GT_CONFLICT" + ) + + hc_int_pixels = int( + row.get( + "highconf_interior_wrong_pixels", + 0, + ) + ) + + if ( + safe_float( + row.get( + "highconf_interior_disagree_frac" + ) + ) + >= float( + min_highconf_interior_frac + ) + and hc_int_pixels + >= int( + min_conflict_pixels + ) + ): + reasons.append( + "HIGH_CONF_INTERIOR_CONFLICT" + ) + + unsafe_pixels = int( + row.get( + "highconf_unsafe_nav_interior_pixels", + 0, + ) + ) + + if ( + safe_float( + row.get( + "highconf_unsafe_nav_interior_frac_all" + ) + ) + >= float( + min_highconf_unsafe_frac + ) + and unsafe_pixels + >= max( + 16, + int( + min_conflict_pixels + // 4 + ), + ) + ): + reasons.append( + "HIGH_CONF_UNSAFE_NAV" + ) + + if ( + int( + row.get( + "status_mismatch", + 0, + ) + or 0 + ) + == 1 + ): + if ( + float( + status_confidence + ) + >= float( + min_status_confidence + ) + ): + reasons.append( + "HIGH_CONF_STATUS_MISMATCH" + ) + else: + reasons.append( + "STATUS_MISMATCH" + ) + + if ( + int( + row.get( + "group_mismatch", + 0, + ) + or 0 + ) + == 1 + and safe_float( + row.get( + "semantic_disagree_frac" + ) + ) + >= float( + min_group_disagree_frac + ) + ): + reasons.append( + "GROUP_COMPOSITION_MISMATCH" + ) + + return list( + dict.fromkeys( + reasons + ) + ) + + +# ============================================================================= +# Split manifest +# ============================================================================= + +def load_split_manifest( + resolution_root: Path, +) -> Dict[ + Tuple[ + str, + str, + ], + str, +]: + path = ( + resolution_root + / "split_manifest.csv" + ) + + if not path.is_file(): + return {} + + result = {} + + with path.open( + "r", + newline="", + encoding="utf-8-sig", + ) as f: + reader = csv.DictReader( + f + ) + + for row in reader: + group = str( + row.get( + "group", + "", + ) + ) + base = str( + row.get( + "base", + "", + ) + ) + split = str( + row.get( + "split", + "", + ) + ) + + if ( + group + and base + and split + ): + result[ + ( + group, + base, + ) + ] = split + + return result + + +# ============================================================================= +# Original-source resolution +# ============================================================================= + +LEGACY_BASE_SUFFIXES = ( + "_rgb", + "_image", + "_img", + "_frame", + "_mask", + "_masks", + "_seg", + "_segment", + "_segmentacao", + "_label", + "_labels", +) + + +def canonical_base( + stem: str, +) -> str: + out = str( + stem + ) + + changed = True + + while changed: + changed = False + low = out.lower() + + for suffix in LEGACY_BASE_SUFFIXES: + if ( + low.endswith( + suffix + ) + and len( + out + ) + > len( + suffix + ) + ): + out = out[ + :-len( + suffix + ) + ] + changed = True + break + + return out + + +def index_folder_canonical( + folder: Path, + exts: set[ + str + ], +) -> Dict[ + str, + List[ + Path + ], +]: + out: Dict[ + str, + List[ + Path + ], + ] = defaultdict( + list + ) + + if not folder.is_dir(): + return out + + for p in folder.iterdir(): + if ( + p.is_file() + and p.suffix.lower() + in exts + ): + out[ + canonical_base( + p.stem + ) + ].append( + p + ) + + return out + + +def resolve_path_value( + raw_value, + *, + dataset_root: Path, + label_path: Path, +) -> Optional[ + Path +]: + if raw_value is None: + return None + + txt = str( + raw_value + ).strip() + + if not txt: + return None + + raw = Path( + txt + ) + + candidates = [] + + if raw.is_absolute(): + candidates.append( + raw + ) + else: + candidates.extend([ + dataset_root + / raw, + label_path.parent + / raw, + label_path.parent.parent.parent.parent + / raw, + ]) + + for p in candidates: + try: + rp = p.resolve() + except Exception: + rp = p + + if rp.is_file(): + return rp + + return None + + +def resolve_original_source( + *, + sample, + dataset_root: Path, +) -> dict: + """ + Prioridade: + 1) provenance source_image/source_mask/source_label do JSON normalizado; + 2) busca em dataset/original/group/ usando base canônica. + + Não escolhe colisão silenciosamente. + """ + normalized_label = ( + Path( + sample.label + ) + if sample.label is not None + else None + ) + + if ( + normalized_label + is not None + and normalized_label.is_file() + ): + data = load_json( + normalized_label + ) + + source_image = resolve_path_value( + data.get( + "source_image" + ), + dataset_root=dataset_root, + label_path=normalized_label, + ) + + source_mask = resolve_path_value( + data.get( + "source_mask" + ), + dataset_root=dataset_root, + label_path=normalized_label, + ) + + source_label = resolve_path_value( + data.get( + "source_label" + ), + dataset_root=dataset_root, + label_path=normalized_label, + ) + + if ( + source_image is not None + and source_mask is not None + ): + return { + "mode": "normalized_provenance", + "image": source_image, + "mask": source_mask, + "label": ( + source_label + if source_label is not None + else normalized_label + ), + } + + # Fallback deterministic to original/group. + group = str( + sample.group + ) + + group_root = ( + dataset_root + / "original" + / "group" + / group + ) + + images = index_folder_canonical( + group_root + / "images", + IMAGE_EXTS, + ) + + masks = index_folder_canonical( + group_root + / "masks", + IMAGE_EXTS, + ) + + labels = index_folder_canonical( + group_root + / "labels", + { + ".json" + }, + ) + + base = canonical_base( + sample.base + ) + + im = images.get( + base, + [], + ) + mk = masks.get( + base, + [], + ) + lb = labels.get( + base, + [], + ) + + if ( + len( + im + ) + != 1 + or len( + mk + ) + != 1 + ): + raise RuntimeError( + f"Fonte original ambígua/ausente para " + f"{sample.group}/{sample.base}: " + f"images={len(im)} masks={len(mk)} labels={len(lb)}" + ) + + label = ( + lb[ + 0 + ] + if len( + lb + ) + == 1 + else normalized_label + ) + + return { + "mode": "original_group_canonical", + "image": im[ + 0 + ], + "mask": mk[ + 0 + ], + "label": label, + } + + +# ============================================================================= +# Physical review tree +# ============================================================================= + +def clear_generated_review( + review_group_root: Path, +) -> None: + """ + Limpa artefatos regeneráveis e PRESERVA final_masks/. + """ + if not review_group_root.exists(): + review_group_root.mkdir( + parents=True, + exist_ok=True, + ) + return + + for group_dir in review_group_root.iterdir(): + if not group_dir.is_dir(): + continue + + for name in ( + "original_images", + "original_masks", + "original_labels", + "predictions", + "panels", + ): + p = ( + group_dir + / name + ) + + if p.exists(): + shutil.rmtree( + p + ) + + order = ( + group_dir + / "review_order.csv" + ) + + if order.exists(): + order.unlink() + + +def write_png( + path: Path, + image: np.ndarray, +) -> None: + path.parent.mkdir( + parents=True, + exist_ok=True, + ) + + ok = cv2.imwrite( + str( + path + ), + image, + [ + cv2.IMWRITE_PNG_COMPRESSION, + 2, + ], + ) + + if not ok: + raise RuntimeError( + f"Falha ao salvar PNG: {path}" + ) + + +def ensure_bgr3( + image: np.ndarray, +) -> np.ndarray: + if image.ndim == 2: + return cv2.cvtColor( + image, + cv2.COLOR_GRAY2BGR, + ) + + if ( + image.ndim + == 3 + and image.shape[ + 2 + ] + == 4 + ): + return cv2.cvtColor( + image, + cv2.COLOR_BGRA2BGR, + ) + + if ( + image.ndim + == 3 + and image.shape[ + 2 + ] + >= 3 + ): + return image[ + :, + :, + :3, + ] + + raise RuntimeError( + f"Imagem inválida shape={image.shape}" + ) + + +def error_map_bgr( + *, + gt: np.ndarray, + pred: np.ndarray, + nav_id: int, + non_nav_id: int, + ignore_id: int, +) -> np.ndarray: + if pred.shape != gt.shape: + pred = cv2.resize( + pred.astype( + np.uint8 + ), + ( + gt.shape[ + 1 + ], + gt.shape[ + 0 + ], + ), + interpolation=cv2.INTER_NEAREST, + ) + + out = np.zeros( + ( + gt.shape[ + 0 + ], + gt.shape[ + 1 + ], + 3, + ), + dtype=np.uint8, + ) + + out[ + : + ] = ( + 60, + 60, + 60, + ) + + valid = ( + gt + != int( + ignore_id + ) + ) + + correct = ( + valid + & ( + gt + == pred + ) + ) + + correct_nav = ( + correct + & ( + gt + == nav_id + ) + ) + + unsafe = ( + valid + & ( + gt + == non_nav_id + ) + & ( + pred + == nav_id + ) + ) + + blocked = ( + valid + & ( + gt + == nav_id + ) + & ( + pred + == non_nav_id + ) + ) + + out[ + correct + ] = ( + 90, + 90, + 90, + ) + out[ + correct_nav + ] = ( + 60, + 155, + 60, + ) + out[ + blocked + ] = ( + 0, + 210, + 255, + ) + out[ + unsafe + ] = ( + 0, + 0, + 255, + ) + out[ + ~valid + ] = ( + 180, + 0, + 180, + ) + + return out + + +def add_title( + image: np.ndarray, + title: str, +) -> np.ndarray: + out = image.copy() + + cv2.rectangle( + out, + ( + 0, + 0, + ), + ( + out.shape[ + 1 + ], + 38, + ), + ( + 18, + 18, + 18, + ), + -1, + ) + + cv2.putText( + out, + title, + ( + 10, + 26, + ), + cv2.FONT_HERSHEY_SIMPLEX, + 0.62, + ( + 245, + 245, + 245, + ), + 1, + cv2.LINE_AA, + ) + + return out + + +def letterbox( + image: np.ndarray, + target_w: int, + target_h: int, +) -> np.ndarray: + h, w = image.shape[ + :2 + ] + + scale = min( + target_w + / max( + 1, + w, + ), + target_h + / max( + 1, + h, + ), + ) + + nw = max( + 1, + int( + round( + w + * scale + ) + ), + ) + nh = max( + 1, + int( + round( + h + * scale + ) + ), + ) + + resized = cv2.resize( + image, + ( + nw, + nh, + ), + interpolation=( + cv2.INTER_AREA + if scale + < 1 + else cv2.INTER_LINEAR + ), + ) + + canvas = np.zeros( + ( + target_h, + target_w, + 3, + ), + dtype=np.uint8, + ) + + canvas[ + : + ] = ( + 24, + 24, + 24, + ) + + x0 = ( + target_w + - nw + ) // 2 + y0 = ( + target_h + - nh + ) // 2 + + canvas[ + y0: + y0 + + nh, + x0: + x0 + + nw, + ] = resized + + return canvas + + +def make_review_panel( + *, + original_bgr: np.ndarray, + gt_color_bgr: np.ndarray, + pred_color_bgr: np.ndarray, + error_bgr: np.ndarray, + row: dict, + label_name_by_id: Dict[ + int, + str, + ], + width: int = 1800, +) -> np.ndarray: + h, w = original_bgr.shape[ + :2 + ] + + gt_overlay = cv2.addWeighted( + original_bgr, + 0.58, + gt_color_bgr, + 0.42, + 0.0, + ) + + pred_overlay = cv2.addWeighted( + original_bgr, + 0.58, + pred_color_bgr, + 0.42, + 0.0, + ) + + cols = 4 + tile_w = max( + 280, + int( + width + / cols + ), + ) + tile_h = max( + 230, + int( + tile_w + * h + / max( + 1, + w, + ) + ), + ) + + tiles = [ + ( + "ORIGINAL", + original_bgr, + ), + ( + "GT", + gt_overlay, + ), + ( + "PRED", + pred_overlay, + ), + ( + "ERRO | vermelho = falso navegável", + error_bgr, + ), + ] + + top = np.concatenate( + [ + add_title( + letterbox( + img, + tile_w, + tile_h, + ), + title, + ) + for title, img in tiles + ], + axis=1, + ) + + panel_h = 220 + + info = np.zeros( + ( + panel_h, + top.shape[ + 1 + ], + 3, + ), + dtype=np.uint8, + ) + + info[ + : + ] = ( + 28, + 28, + 28, + ) + + gt_status_id = row.get( + "gt_status_id" + ) + + pred_status_id = row.get( + "pred_status_id" + ) + + gt_status = ( + label_name_by_id.get( + int( + gt_status_id + ), + str( + gt_status_id + ), + ) + if gt_status_id + is not None + else "-" + ) + + pred_status = ( + label_name_by_id.get( + int( + pred_status_id + ), + str( + pred_status_id + ), + ) + if pred_status_id + is not None + else "-" + ) + + lines = [ + ( + f"{row['group']}/{row['base']} | " + f"split={row.get('split_origin','?')} | " + f"score={row['suspicion_pct']:.1f}/100" + ), + ( + f"SEG: navIoU={row['nav_iou']:.4f} " + f"nonNavIoU={row['non_nav_iou']:.4f} " + f"disagree={100.0*row['semantic_disagree_frac']:.2f}% " + f"HC-interior={100.0*row['highconf_interior_disagree_frac']:.3f}%" + ), + ( + f"RISK: unsafe={100.0*row['unsafe_nav_frac_gt_nonnav']:.3f}% " + f"HC-unsafe-interior={100.0*row['highconf_unsafe_nav_interior_frac_all']:.3f}%" + ), + ( + f"STATUS: GT={gt_status} | PRED={pred_status} " + f"conf={row['status_confidence']:.3f} " + f"GTprob={safe_float(row.get('gt_status_probability'),0.0):.3f}" + ), + ( + f"GROUP: GT={row['group']} | PRED={row['pred_group']} " + f"| reasons={row.get('review_reasons','')}" + ), + ] + + y = 34 + + for line in lines: + cv2.putText( + info, + line, + ( + 14, + y, + ), + cv2.FONT_HERSHEY_SIMPLEX, + 0.62, + ( + 240, + 240, + 240, + ), + 1, + cv2.LINE_AA, + ) + + y += 37 + + return np.concatenate( + [ + top, + info, + ], + axis=0, + ) + + +# ============================================================================= +# Inference +# ============================================================================= + +@torch.inference_mode() +def infer_audit( + *, + base_module, + base_model, + status_head, + contract, + image_bgr: np.ndarray, + device: torch.device, + use_amp: bool, +) -> dict: + x, geometry = base_module.prepare_input( + image_bgr, + contract, + device, + ) + + if ( + device.type + == "cuda" + ): + torch.cuda.synchronize() + + t0 = time.perf_counter() + + amp_enabled = ( + bool( + use_amp + ) + and device.type + == "cuda" + ) + + with torch.autocast( + device_type=( + "cuda" + if device.type + == "cuda" + else "cpu" + ), + dtype=torch.float16, + enabled=amp_enabled, + ): + out = base_model( + pixel_values=x, + output_hidden_states=True, + return_dict=True, + ) + + seg_logits_native = out.logits + feat = out.hidden_states[ + -1 + ] + + status_logits = status_head( + feat, + seg_logits_native, + ) + + W, H = contract.resolution_wh + + seg_logits_full = F.interpolate( + seg_logits_native, + size=( + H, + W, + ), + mode="bilinear", + align_corners=False, + ) + + seg_probs = torch.softmax( + seg_logits_full.float(), + dim=1, + ) + + pred = torch.argmax( + seg_probs, + dim=1, + ) + + status_probs = torch.softmax( + status_logits.float(), + dim=1, + ) + + if ( + device.type + == "cuda" + ): + torch.cuda.synchronize() + + infer_ms = ( + time.perf_counter() + - t0 + ) * 1000.0 + + return { + "pred": ( + pred[ + 0 + ] + .detach() + .cpu() + .numpy() + .astype( + np.uint8 + ) + ), + "seg_probs": ( + seg_probs[ + 0 + ] + .detach() + .cpu() + .numpy() + .astype( + np.float32 + ) + ), + "status_probs": ( + status_probs[ + 0 + ] + .detach() + .cpu() + .numpy() + .astype( + np.float32 + ) + ), + "geometry": geometry, + "infer_ms": float( + infer_ms + ), + } + + +# ============================================================================= +# Candidate export +# ============================================================================= + +def save_review_item( + *, + row: dict, + sample, + pred_ids: np.ndarray, + gt_ids: np.ndarray, + source: dict, + review_group_root: Path, + base_module, + classes, + contract, + save_panels: bool, +) -> dict: + group = str( + sample.group + ) + + group_root = ( + review_group_root + / group + ) + + dirs = { + "images": group_root + / "original_images", + "masks": group_root + / "original_masks", + "labels": group_root + / "original_labels", + "predictions": group_root + / "predictions", + "panels": group_root + / "panels", + "final_masks": group_root + / "final_masks", + } + + for name, folder in dirs.items(): + if ( + name + == "panels" + and not save_panels + ): + continue + + folder.mkdir( + parents=True, + exist_ok=True, + ) + + src_image = Path( + source[ + "image" + ] + ) + src_mask = Path( + source[ + "mask" + ] + ) + src_label = ( + Path( + source[ + "label" + ] + ) + if source.get( + "label" + ) + is not None + else None + ) + + original_bgr = cv2.imread( + str( + src_image + ), + cv2.IMREAD_COLOR, + ) + + if original_bgr is None: + raise RuntimeError( + f"Falha imagem original: {src_image}" + ) + + original_mask_native = cv2.imread( + str( + src_mask + ), + cv2.IMREAD_UNCHANGED, + ) + + if original_mask_native is None: + raise RuntimeError( + f"Falha mask original: {src_mask}" + ) + + original_mask_bgr = ensure_bgr3( + original_mask_native + ) + + oh, ow = original_bgr.shape[ + :2 + ] + + mh, mw = original_mask_bgr.shape[ + :2 + ] + + if ( + oh, + ow, + ) != ( + mh, + mw, + ): + raise RuntimeError( + f"Imagem/mask original geometria diferente: " + f"{src_image}={ow}x{oh}, {src_mask}={mw}x{mh}" + ) + + # O normalize do corredor só faz resize mantendo aspect ratio. + pred_original_ids = cv2.resize( + pred_ids.astype( + np.uint8 + ), + ( + ow, + oh, + ), + interpolation=cv2.INTER_NEAREST, + ) + + gt_original_ids = cv2.resize( + gt_ids.astype( + np.uint8 + ), + ( + ow, + oh, + ), + interpolation=cv2.INTER_NEAREST, + ) + + pred_color_bgr = base_module.colorize_ids( + pred_original_ids, + classes, + ) + + gt_color_bgr = base_module.colorize_ids( + gt_original_ids, + classes, + ) + + error_bgr = error_map_bgr( + gt=gt_original_ids, + pred=pred_original_ids, + nav_id=contract.nav_class_id, + non_nav_id=contract.non_nav_class_id, + ignore_id=IGNORE_INDEX, + ) + + image_out = ( + dirs[ + "images" + ] + / f"{sample.base}{src_image.suffix.lower()}" + ) + + mask_out = ( + dirs[ + "masks" + ] + / f"{sample.base}{src_mask.suffix.lower()}" + ) + + pred_out = ( + dirs[ + "predictions" + ] + / f"{sample.base}.png" + ) + + shutil.copy2( + src_image, + image_out, + ) + + shutil.copy2( + src_mask, + mask_out, + ) + + label_out = None + + if ( + src_label is not None + and src_label.is_file() + ): + label_out = ( + dirs[ + "labels" + ] + / f"{sample.base}{src_label.suffix.lower()}" + ) + + shutil.copy2( + src_label, + label_out, + ) + + write_png( + pred_out, + pred_color_bgr, + ) + + panel_out = None + + if save_panels: + panel = make_review_panel( + original_bgr=original_bgr, + gt_color_bgr=gt_color_bgr, + pred_color_bgr=pred_color_bgr, + error_bgr=error_bgr, + row=row, + label_name_by_id=contract.label_name_by_id, + ) + + panel_out = ( + dirs[ + "panels" + ] + / f"{sample.base}.jpg" + ) + + ok = cv2.imwrite( + str( + panel_out + ), + panel, + [ + cv2.IMWRITE_JPEG_QUALITY, + 92, + ], + ) + + if not ok: + raise RuntimeError( + f"Falha panel: {panel_out}" + ) + + return { + "group": group, + "base": str( + sample.base + ), + "source_mode": str( + source[ + "mode" + ] + ), + "source_image": str( + src_image + ), + "source_mask": str( + src_mask + ), + "source_label": ( + str( + src_label + ) + if src_label + is not None + else None + ), + "review_image": str( + image_out + ), + "review_mask": str( + mask_out + ), + "review_label": ( + str( + label_out + ) + if label_out + is not None + else None + ), + "review_prediction": str( + pred_out + ), + "review_panel": ( + str( + panel_out + ) + if panel_out + is not None + else None + ), + "width": int( + ow + ), + "height": int( + oh + ), + "prediction_geometry": ( + "normalized_model_mask resized with INTER_NEAREST " + "to original annotation geometry" + ), + } + + +# ============================================================================= +# Summary helpers +# ============================================================================= + +def summarize_rows( + rows: Sequence[ + dict + ], + selected_keys: set[ + Tuple[ + str, + str, + ] + ], +) -> dict: + if not rows: + return {} + + scores = np.asarray( + [ + safe_float( + r.get( + "review_score" + ) + ) + for r in rows + ], + dtype=np.float64, + ) + + disag = np.asarray( + [ + safe_float( + r.get( + "semantic_disagree_frac" + ) + ) + for r in rows + ], + dtype=np.float64, + ) + + hc = np.asarray( + [ + safe_float( + r.get( + "highconf_disagree_frac" + ) + ) + for r in rows + ], + dtype=np.float64, + ) + + hc_int = np.asarray( + [ + safe_float( + r.get( + "highconf_interior_disagree_frac" + ) + ) + for r in rows + ], + dtype=np.float64, + ) + + unsafe = np.asarray( + [ + safe_float( + r.get( + "highconf_unsafe_nav_interior_frac_all" + ) + ) + for r in rows + ], + dtype=np.float64, + ) + + selected = sum( + ( + str( + r[ + "group" + ] + ), + str( + r[ + "base" + ] + ), + ) + in selected_keys + for r in rows + ) + + return { + "samples": len( + rows + ), + "selected": int( + selected + ), + "selected_pct": float( + 100.0 + * selected + / max( + 1, + len( + rows + ), + ) + ), + "review_score_mean": float( + scores.mean() + ), + "review_score_p95": float( + np.percentile( + scores, + 95, + ) + ), + "review_score_max": float( + scores.max() + ), + "semantic_disagree_mean": float( + disag.mean() + ), + "semantic_disagree_p95": float( + np.percentile( + disag, + 95, + ) + ), + "highconf_disagree_mean": float( + hc.mean() + ), + "highconf_interior_mean": float( + hc_int.mean() + ), + "highconf_unsafe_interior_mean": float( + unsafe.mean() + ), + "status_mismatch_count": int( + sum( + int( + r.get( + "status_mismatch", + 0, + ) + or 0 + ) + for r in rows + ) + ), + } + + +# ============================================================================= +# Main +# ============================================================================= + +def main(): + ap = argparse.ArgumentParser( + description=( + "Auditoria modelo x GT do dataset frontal de corredor." + ) + ) + + ap.add_argument( + "--config", + default="config.json", + ) + + ap.add_argument( + "--test_script", + default=None, + help=( + "Default: test_corridor_agri_v2.py ao lado deste script." + ), + ) + + ap.add_argument( + "--ckpt", + default=None, + help="Checkpoint .pt explícito.", + ) + + ap.add_argument( + "--checkpoint", + default="operational", + choices=[ + "operational", + "nav", + "status", + "last", + ], + ) + + ap.add_argument( + "--root", + default=None, + help=( + "Default: dataset/x/group." + ), + ) + + ap.add_argument( + "--device", + default="cuda", + choices=[ + "cuda", + "cpu", + ], + ) + + ap.add_argument( + "--no_amp", + action="store_true", + ) + + ap.add_argument( + "--out_root", + default="dataset/revisao_corridor", + ) + + ap.add_argument( + "--run_name", + default=None, + ) + + ap.add_argument( + "--report_only", + action="store_true", + ) + + ap.add_argument( + "--save_panels", + action="store_true", + ) + + ap.add_argument( + "--clear_review", + action="store_true", + help=( + "Limpa artefatos regeneráveis, preservando final_masks." + ), + ) + + ap.add_argument( + "--clear_review_all", + action="store_true", + help=( + "PERIGOSO: apaga toda a árvore física de revisão." + ), + ) + + ap.add_argument( + "--export_mode", + default="flagged", + choices=[ + "flagged", + "score", + "all", + ], + ) + + ap.add_argument( + "--min_suspicion_pct", + type=float, + default=65.0, + help=( + "Usado em --export_mode score. Ranking heurístico 0..100." + ), + ) + + ap.add_argument( + "--top_k_per_stratum", + type=int, + default=0, + help=( + "Se >0, seleciona os K scores maiores de cada grupo x GT status." + ), + ) + + ap.add_argument( + "--max_review_per_stratum", + type=int, + default=0, + help="Cap final por grupo x status. 0=sem limite.", + ) + + ap.add_argument( + "--max_samples", + type=int, + default=0, + help="Debug. 0=todas.", + ) + + ap.add_argument( + "--gc_every", + type=int, + default=50, + ) + + # Semantic mining. + ap.add_argument( + "--high_confidence", + type=float, + default=0.90, + ) + + ap.add_argument( + "--uncertainty_margin", + type=float, + default=0.10, + ) + + ap.add_argument( + "--boundary_radius", + type=int, + default=4, + help="Raio px da banda de fronteira do GT.", + ) + + ap.add_argument( + "--composition_min_fraction", + type=float, + default=0.005, + help=( + "0.005=0.5%%. Menor que isso é tratado como composição pura." + ), + ) + + # Technical flags. + ap.add_argument( + "--min_disagree_frac", + type=float, + default=0.05, + ) + + ap.add_argument( + "--min_highconf_disagree_frac", + type=float, + default=0.002, + ) + + ap.add_argument( + "--min_highconf_interior_frac", + type=float, + default=0.001, + ) + + ap.add_argument( + "--min_highconf_unsafe_frac", + type=float, + default=0.0005, + ) + + ap.add_argument( + "--min_conflict_pixels", + type=int, + default=150, + ) + + ap.add_argument( + "--min_status_confidence", + type=float, + default=0.80, + ) + + ap.add_argument( + "--min_group_disagree_frac", + type=float, + default=0.05, + ) + + args = ap.parse_args() + + if ( + args.clear_review + and args.clear_review_all + ): + ap.error( + "Use apenas --clear_review OU --clear_review_all." + ) + + if not ( + 0.0 + <= float( + args.min_suspicion_pct + ) + <= 100.0 + ): + ap.error( + "--min_suspicion_pct deve ficar entre 0 e 100." + ) + + if not ( + 0.0 + < float( + args.high_confidence + ) + <= 1.0 + ): + ap.error( + "--high_confidence inválido." + ) + + # ------------------------------------------------------------------------- + # Shared tester + # ------------------------------------------------------------------------- + + here = Path( + __file__ + ).resolve() + + test_script = ( + Path( + args.test_script + ).resolve() + if args.test_script + else here.with_name( + "test_corridor_agri_v2.py" + ) + ) + + base = load_test_module( + test_script + ) + + config_path = Path( + args.config + ).resolve() + + if not config_path.is_file(): + raise FileNotFoundError( + config_path + ) + + config = base.safe_json_load( + config_path + ) + + project_root = base.find_project_root( + config_path + ) + + dataset_root = ( + project_root + / "dataset" + ) + + # ------------------------------------------------------------------------- + # Checkpoint + # ------------------------------------------------------------------------- + + if args.ckpt: + checkpoint_path = Path( + args.ckpt + ).resolve() + else: + checkpoint_path = base.default_checkpoint_path( + project_root, + config, + args.checkpoint, + ) + + use_cuda = ( + args.device + == "cuda" + and torch.cuda.is_available() + ) + + if ( + args.device + == "cuda" + and not use_cuda + ): + print( + "[WARN] CUDA indisponível; usando CPU." + ) + + device = torch.device( + "cuda" + if use_cuda + else "cpu" + ) + + ( + base_model, + status_head, + contract, + checkpoint_raw, + ) = base.build_runtime_from_checkpoint( + checkpoint_path, + config, + device, + ) + + W, H = contract.resolution_wh + + resolution_root = ( + dataset_root + / f"{W}x{H}" + ) + + root = ( + Path( + args.root + ).resolve() + if args.root + else ( + resolution_root + / "group" + ) + ) + + if not root.is_dir(): + raise FileNotFoundError( + f"Dataset auditável não encontrado: {root}" + ) + + classes = base.load_labelmap( + dataset_root + / "labelmap.txt" + ) + + samples, source_kind = base.discover_samples( + root + ) + + samples = [ + s + for s in samples + if ( + s.mask + is not None + and s.label + is not None + ) + ] + + if args.max_samples > 0: + samples = samples[ + :int( + args.max_samples + ) + ] + + if not samples: + raise RuntimeError( + "Nenhuma tripleta image+mask+label encontrada." + ) + + split_origin_map = load_split_manifest( + resolution_root + ) + + # ------------------------------------------------------------------------- + # Output + # ------------------------------------------------------------------------- + + out_root = Path( + args.out_root + ) + + if not out_root.is_absolute(): + out_root = ( + project_root + / out_root + ) + + out_root = out_root.resolve() + + review_group_root = ( + out_root + / "group" + ) + + run_name = ( + args.run_name + or checkpoint_path.stem + ) + + report_dir = ( + out_root + / "reports" + / run_name + ) + + report_dir.mkdir( + parents=True, + exist_ok=True, + ) + + if not args.report_only: + if args.clear_review_all: + print( + "[WARN] Limpando revisão COMPLETA, inclusive final_masks." + ) + clear_dir( + review_group_root + ) + elif args.clear_review: + print( + "[INFO] Limpando artefatos regeneráveis; final_masks preservada." + ) + clear_generated_review( + review_group_root + ) + + # ------------------------------------------------------------------------- + # Print contract + # ------------------------------------------------------------------------- + + print("=" * 94) + print( + "AGROBOT | CORRIDOR DATASET REVIEW / MODEL MINING" + ) + print("=" * 94) + print( + f"Reviewer : {REVIEWER_VERSION}" + ) + print( + f"Tester base : {test_script}" + ) + print( + f"Dataset root : {root}" + ) + print( + f"Source kind : {source_kind}" + ) + print( + f"Samples : {len(samples)}" + ) + print( + f"Checkpoint : {checkpoint_path}" + ) + print( + f"Epoch : {contract.checkpoint_epoch}" + ) + print( + f"Resolution : {W}x{H}" + ) + print( + f"Device : {device}" + ) + print( + f"AMP : {not args.no_amp and device.type == 'cuda'}" + ) + print( + f"High confidence : {args.high_confidence:.3f}" + ) + print( + f"Boundary radius : {args.boundary_radius}px" + ) + print( + f"Export mode : {args.export_mode}" + ) + print( + f"Report only : {args.report_only}" + ) + print("=" * 94) + + # ------------------------------------------------------------------------- + # Audit pass + # ------------------------------------------------------------------------- + + num_seg_classes = len( + contract.seg_id2label + ) + + num_status_classes = len( + contract.label_name_by_id + ) + + cm_seg_total = np.zeros( + ( + num_seg_classes, + num_seg_classes, + ), + dtype=np.int64, + ) + + cm_status_total = np.zeros( + ( + num_status_classes, + num_status_classes, + ), + dtype=np.int64, + ) + + rows: List[ + dict + ] = [] + + infer_ms_sum = 0.0 + + t0_all = time.perf_counter() + + for i, sample in enumerate( + samples, + 1, + ): + image_bgr = cv2.imread( + str( + sample.image + ), + cv2.IMREAD_COLOR, + ) + + if image_bgr is None: + raise RuntimeError( + f"Imagem inválida: {sample.image}" + ) + + gt = base.load_gt_mask( + sample.mask, + classes, + ) + + if gt is None: + raise RuntimeError( + f"GT ausente: {sample.mask}" + ) + + ( + gt_status_id, + gt_status_name, + ) = base.read_gt_label( + sample, + contract.label_name_by_id, + ) + + infer = infer_audit( + base_module=base, + base_model=base_model, + status_head=status_head, + contract=contract, + image_bgr=image_bgr, + device=device, + use_amp=not args.no_amp, + ) + + infer_ms_sum += float( + infer[ + "infer_ms" + ] + ) + + # Dataset normalizado deveria já estar HxW. + # Mesmo assim deixa robusto. + pred = infer[ + "pred" + ] + + seg_probs = infer[ + "seg_probs" + ] + + if gt.shape != pred.shape: + gt_eval = cv2.resize( + gt.astype( + np.uint8 + ), + ( + pred.shape[ + 1 + ], + pred.shape[ + 0 + ], + ), + interpolation=cv2.INTER_NEAREST, + ) + else: + gt_eval = gt + + cm_sample = confusion_matrix_np( + pred, + gt_eval, + num_seg_classes, + IGNORE_INDEX, + ) + + cm_seg_total += cm_sample + + metric_sample = binary_metrics_from_cm( + cm_sample, + contract.nav_class_id, + contract.non_nav_class_id, + ) + + sem = semantic_sample_signals( + pred=pred, + probs=seg_probs, + gt=gt_eval, + nav_id=contract.nav_class_id, + non_nav_id=contract.non_nav_class_id, + ignore_id=IGNORE_INDEX, + high_confidence=float( + args.high_confidence + ), + uncertainty_margin=float( + args.uncertainty_margin + ), + boundary_radius=int( + args.boundary_radius + ), + ) + + status = status_sample_signals( + infer[ + "status_probs" + ], + gt_status_id, + ) + + pred_status_id = int( + status[ + "pred_status_id" + ] + ) + + if ( + gt_status_id + is not None + and 0 + <= int( + gt_status_id + ) + < num_status_classes + ): + cm_status_total[ + int( + gt_status_id + ), + pred_status_id, + ] += 1 + + pred_group, pred_nav_frac = derive_composition_group( + pred, + contract.nav_class_id, + contract.non_nav_class_id, + float( + args.composition_min_fraction + ), + ) + + gt_group = str( + sample.group + ) + + group_mismatch = ( + pred_group + != gt_group + ) + + split_origin = split_origin_map.get( + ( + gt_group, + str( + sample.base + ), + ), + ( + "unknown" + if ( + "split" + not in str( + root + ).lower() + ) + else root.name + ), + ) + + row = { + "index": i - 1, + "group": gt_group, + "base": str( + sample.base + ), + "split_origin": split_origin, + + "image_path": str( + sample.image + ), + "mask_path": str( + sample.mask + ), + "label_path": str( + sample.label + ), + + "inference_ms": float( + infer[ + "infer_ms" + ] + ), + + **metric_sample, + **sem, + + "gt_status_id": ( + int( + gt_status_id + ) + if gt_status_id + is not None + else None + ), + "gt_status_name": gt_status_name, + + **status, + + "pred_status_name": contract.label_name_by_id.get( + pred_status_id, + str( + pred_status_id + ), + ), + + "pred_group": pred_group, + "pred_nav_frac": float( + pred_nav_frac + ), + "group_mismatch": int( + group_mismatch + ), + } + + score, score_parts = compute_review_score( + row + ) + + row[ + "review_score" + ] = float( + score + ) + row[ + "suspicion_pct" + ] = float( + score + * 100.0 + ) + + row.update( + score_parts + ) + + reasons = build_reasons( + row, + min_disagree_frac=float( + args.min_disagree_frac + ), + min_highconf_disagree_frac=float( + args.min_highconf_disagree_frac + ), + min_highconf_interior_frac=float( + args.min_highconf_interior_frac + ), + min_highconf_unsafe_frac=float( + args.min_highconf_unsafe_frac + ), + min_conflict_pixels=int( + args.min_conflict_pixels + ), + status_confidence=float( + row[ + "status_confidence" + ] + ), + min_status_confidence=float( + args.min_status_confidence + ), + min_group_disagree_frac=float( + args.min_group_disagree_frac + ), + ) + + row[ + "review_reasons" + ] = ";".join( + reasons + ) + + row[ + "flagged" + ] = int( + bool( + reasons + ) + ) + + row[ + "stratum" + ] = ( + f"{gt_group}" + f"__" + f"{gt_status_name or 'unknown'}" + ) + + rows.append( + row + ) + + del ( + image_bgr, + gt, + gt_eval, + pred, + seg_probs, + infer, + ) + + if ( + args.gc_every + > 0 + and i + % int( + args.gc_every + ) + == 0 + ): + gc.collect() + + if ( + device.type + == "cuda" + ): + torch.cuda.empty_cache() + + if ( + i + % 50 + == 0 + or i + == len( + samples + ) + ): + flagged_now = sum( + int( + r[ + "flagged" + ] + ) + for r in rows + ) + + print( + f"[{i:5d}/{len(samples):5d}] " + f"flagged={flagged_now:4d} " + f"avg_inf={infer_ms_sum / i:.1f}ms " + f"last_score={score*100.0:5.1f} " + f"{sample.group}/{sample.base}" + ) + + elapsed = ( + time.perf_counter() + - t0_all + ) + + # ------------------------------------------------------------------------- + # Candidate selection + # ------------------------------------------------------------------------- + + by_stratum: Dict[ + str, + List[ + dict + ], + ] = defaultdict( + list + ) + + for row in rows: + by_stratum[ + str( + row[ + "stratum" + ] + ) + ].append( + row + ) + + selected_rows: List[ + dict + ] = [] + + for stratum, srows in sorted( + by_stratum.items() + ): + ordered = sorted( + srows, + key=lambda r: float( + r[ + "review_score" + ] + ), + reverse=True, + ) + + if args.top_k_per_stratum > 0: + chosen = ordered[ + :int( + args.top_k_per_stratum + ) + ] + + for row in chosen: + if not row[ + "review_reasons" + ]: + row[ + "review_reasons" + ] = "TOP_K_STRATUM" + row[ + "flagged" + ] = 1 + + elif ( + args.export_mode + == "all" + ): + chosen = ordered + + for row in chosen: + if not row[ + "review_reasons" + ]: + row[ + "review_reasons" + ] = "EXPORT_ALL" + + elif ( + args.export_mode + == "score" + ): + thr = ( + float( + args.min_suspicion_pct + ) + / 100.0 + ) + + chosen = [ + row + for row in ordered + if float( + row[ + "review_score" + ] + ) + >= thr + ] + + for row in chosen: + if not row[ + "review_reasons" + ]: + row[ + "review_reasons" + ] = "SUSPICION_SCORE" + + else: + chosen = [ + row + for row in ordered + if int( + row[ + "flagged" + ] + ) + == 1 + ] + + if ( + args.max_review_per_stratum + > 0 + ): + chosen = chosen[ + :int( + args.max_review_per_stratum + ) + ] + + for rank, row in enumerate( + chosen, + 1, + ): + row[ + "review_rank_stratum" + ] = rank + + selected_rows.extend( + chosen + ) + + selected_rows = sorted( + selected_rows, + key=lambda r: ( + -float( + r[ + "review_score" + ] + ), + str( + r[ + "stratum" + ] + ), + natural_key( + r[ + "base" + ] + ), + ), + ) + + for rank, row in enumerate( + selected_rows, + 1, + ): + row[ + "review_rank_global" + ] = rank + + selected_keys = { + ( + str( + r[ + "group" + ] + ), + str( + r[ + "base" + ] + ), + ) + for r in selected_rows + } + + # ------------------------------------------------------------------------- + # Physical export, memory-safe second pass + # ------------------------------------------------------------------------- + + source_manifest_rows = [] + + if not args.report_only: + print() + print( + f"[EXPORT] candidatos físicos: {len(selected_rows)}" + ) + + sample_by_key = { + ( + str( + s.group + ), + str( + s.base + ), + ): s + for s in samples + } + + selected_by_group = defaultdict( + list + ) + + for row in selected_rows: + selected_by_group[ + str( + row[ + "group" + ] + ) + ].append( + row + ) + + for j, row in enumerate( + selected_rows, + 1, + ): + key = ( + str( + row[ + "group" + ] + ), + str( + row[ + "base" + ] + ), + ) + + sample = sample_by_key[ + key + ] + + try: + source = resolve_original_source( + sample=sample, + dataset_root=dataset_root, + ) + + image_bgr = cv2.imread( + str( + sample.image + ), + cv2.IMREAD_COLOR, + ) + + if image_bgr is None: + raise RuntimeError( + sample.image + ) + + gt_ids = base.load_gt_mask( + sample.mask, + classes, + ) + + if gt_ids is None: + raise RuntimeError( + sample.mask + ) + + infer = infer_audit( + base_module=base, + base_model=base_model, + status_head=status_head, + contract=contract, + image_bgr=image_bgr, + device=device, + use_amp=not args.no_amp, + ) + + pred_ids = infer[ + "pred" + ] + + exported = save_review_item( + row=row, + sample=sample, + pred_ids=pred_ids, + gt_ids=gt_ids, + source=source, + review_group_root=review_group_root, + base_module=base, + classes=classes, + contract=contract, + save_panels=bool( + args.save_panels + ), + ) + + row[ + "physical_export_ok" + ] = 1 + row[ + "physical_export_error" + ] = "" + + source_manifest_rows.append({ + **exported, + "review_score": row[ + "review_score" + ], + "suspicion_pct": row[ + "suspicion_pct" + ], + "review_reasons": row[ + "review_reasons" + ], + }) + + except Exception as exc: + row[ + "physical_export_ok" + ] = 0 + row[ + "physical_export_error" + ] = ( + f"{type(exc).__name__}: {exc}" + ) + + print( + f"[EXPORT][WARN] " + f"{row['group']}/{row['base']}: " + f"{row['physical_export_error']}" + ) + + if ( + args.gc_every + > 0 + and j + % int( + args.gc_every + ) + == 0 + ): + gc.collect() + + if ( + device.type + == "cuda" + ): + torch.cuda.empty_cache() + + if ( + j + % 50 + == 0 + or j + == len( + selected_rows + ) + ): + print( + f" exportados {j}/{len(selected_rows)}" + ) + + # review_order.csv por grupo. + for group, grows in selected_by_group.items(): + order_path = ( + review_group_root + / group + / "review_order.csv" + ) + + write_csv( + order_path, + sorted( + grows, + key=lambda r: int( + r.get( + "review_rank_global", + 999999, + ) + ), + ), + ) + + write_csv( + out_root + / "review_source_manifest.csv", + source_manifest_rows, + ) + + # ------------------------------------------------------------------------- + # Reports + # ------------------------------------------------------------------------- + + report_path = ( + report_dir + / "review_report.csv" + ) + + candidates_path = ( + report_dir + / "review_candidates.csv" + ) + + write_csv( + report_path, + rows, + ) + + write_csv( + candidates_path, + selected_rows, + ) + + semantic_labels = [ + contract.seg_id2label[ + i + ] + for i in sorted( + contract.seg_id2label + ) + ] + + status_labels = [ + contract.label_name_by_id[ + i + ] + for i in sorted( + contract.label_name_by_id + ) + ] + + write_cm_csv( + report_dir + / "confusion_semantic.csv", + cm_seg_total, + semantic_labels, + ) + + write_cm_csv( + report_dir + / "confusion_status.csv", + cm_status_total, + status_labels, + ) + + global_seg = binary_metrics_from_cm( + cm_seg_total, + contract.nav_class_id, + contract.non_nav_class_id, + ) + + global_status = status_metrics_from_cm( + cm_status_total + ) + + by_group = defaultdict( + list + ) + by_status = defaultdict( + list + ) + by_split = defaultdict( + list + ) + + for row in rows: + by_group[ + str( + row[ + "group" + ] + ) + ].append( + row + ) + + by_status[ + str( + row.get( + "gt_status_name", + "unknown", + ) + ) + ].append( + row + ) + + by_split[ + str( + row.get( + "split_origin", + "unknown", + ) + ) + ].append( + row + ) + + group_summary = { + key: summarize_rows( + value, + selected_keys, + ) + for key, value in sorted( + by_group.items() + ) + } + + status_summary = { + key: summarize_rows( + value, + selected_keys, + ) + for key, value in sorted( + by_status.items() + ) + } + + split_summary = { + key: summarize_rows( + value, + selected_keys, + ) + for key, value in sorted( + by_split.items() + ) + } + + stratum_summary = { + key: summarize_rows( + value, + selected_keys, + ) + for key, value in sorted( + by_stratum.items() + ) + } + + all_summary = summarize_rows( + rows, + selected_keys, + ) + + source_export_errors = [ + { + "group": row[ + "group" + ], + "base": row[ + "base" + ], + "error": row.get( + "physical_export_error", + "", + ), + } + for row in selected_rows + if int( + row.get( + "physical_export_ok", + 1, + ) + or 0 + ) + == 0 + ] + + summary = { + "schema": REVIEW_SCHEMA, + "reviewer_version": REVIEWER_VERSION, + "created_at": now_iso(), + + "note": ( + "Model-vs-GT disagreement is a review signal, not proof of annotation error. " + "review_score/suspicion_pct are heuristic ranking values, not probabilities." + ), + + "config": str( + config_path + ), + "test_script": str( + test_script + ), + "checkpoint": str( + checkpoint_path + ), + "checkpoint_epoch": contract.checkpoint_epoch, + "trainer_version": contract.trainer_version, + + "dataset_root": str( + root + ), + "resolution_wh": [ + W, + H, + ], + + "samples": len( + rows + ), + "selected_for_review": len( + selected_rows + ), + "selected_pct": float( + 100.0 + * len( + selected_rows + ) + / max( + 1, + len( + rows + ), + ) + ), + + "elapsed_seconds": float( + elapsed + ), + "mean_inference_ms": float( + infer_ms_sum + / max( + 1, + len( + rows + ), + ) + ), + + "selection": { + "export_mode": args.export_mode, + "min_suspicion_pct": float( + args.min_suspicion_pct + ), + "top_k_per_stratum": int( + args.top_k_per_stratum + ), + "max_review_per_stratum": int( + args.max_review_per_stratum + ), + "report_only": bool( + args.report_only + ), + }, + + "thresholds": { + "high_confidence": float( + args.high_confidence + ), + "uncertainty_margin": float( + args.uncertainty_margin + ), + "boundary_radius_px": int( + args.boundary_radius + ), + "composition_min_fraction": float( + args.composition_min_fraction + ), + "min_disagree_frac": float( + args.min_disagree_frac + ), + "min_highconf_disagree_frac": float( + args.min_highconf_disagree_frac + ), + "min_highconf_interior_frac": float( + args.min_highconf_interior_frac + ), + "min_highconf_unsafe_frac": float( + args.min_highconf_unsafe_frac + ), + "min_conflict_pixels": int( + args.min_conflict_pixels + ), + "min_status_confidence": float( + args.min_status_confidence + ), + }, + + "metrics_global": { + "semantic": global_seg, + "status": global_status, + }, + + "annotation_health": all_summary, + + "groups": group_summary, + "statuses": status_summary, + "splits": split_summary, + "strata_group_x_status": stratum_summary, + + "physical_export": { + "enabled": not bool( + args.report_only + ), + "exported": len( + source_manifest_rows + ), + "errors": source_export_errors, + "final_masks_preserved_by_clear_review": True, + }, + + "paths": { + "review_report": str( + report_path + ), + "review_candidates": str( + candidates_path + ), + "confusion_semantic": str( + report_dir + / "confusion_semantic.csv" + ), + "confusion_status": str( + report_dir + / "confusion_status.csv" + ), + "review_group_root": ( + str( + review_group_root + ) + if not args.report_only + else None + ), + "review_source_manifest": ( + str( + out_root + / "review_source_manifest.csv" + ) + if not args.report_only + else None + ), + }, + } + + summary_path = ( + report_dir + / "review_summary.json" + ) + + save_json( + summary_path, + summary, + ) + + readme_path = ( + report_dir + / "README_REVIEW.txt" + ) + + readme_path.write_text( + "\n".join([ + "AGROBOT CORRIDOR DATASET REVIEW / MODEL MINING", + "", + "IMPORTANTE:", + "- Discordancia modelo x GT NAO prova que a anotacao esta errada.", + "- Conflito de alta confianca no interior da regiao e um sinal forte para revisao.", + "- Falso navegavel de alta confianca recebe peso maior por risco operacional.", + "- A segunda cabeca (estado global) e auditada junto da mascara.", + "- review_score e suspicion_pct sao apenas ranking heuristico.", + "- O auditor usa o checkpoint PyTorch porque o ONNX de CAMPO e deliberadamente mask-only.", + "- --clear_review preserva final_masks/.", + "- --clear_review_all apaga tudo e deve ser usado com cuidado.", + "", + f"Checkpoint: {checkpoint_path}", + f"Dataset: {root}", + f"Samples: {len(rows)}", + f"Selecionados: {len(selected_rows)}", + f"Report: {report_path}", + f"Candidates: {candidates_path}", + ]), + encoding="utf-8", + ) + + # ------------------------------------------------------------------------- + # Console + # ------------------------------------------------------------------------- + + print() + print("=" * 94) + print( + "REVISÃO CONCLUÍDA" + ) + print("=" * 94) + print( + f"Samples : {len(rows)}" + ) + print( + f"Selecionados : " + f"{len(selected_rows)} " + f"({100.0 * len(selected_rows) / max(1,len(rows)):.2f}%)" + ) + print( + f"Tempo total : {elapsed:.1f}s" + ) + print( + f"Inferência média : " + f"{infer_ms_sum / max(1,len(rows)):.2f} ms" + ) + print( + f"Semantic acc/mIoU : " + f"{global_seg['acc']:.4f} / " + f"{global_seg['miou']:.4f}" + ) + print( + f"Nav IoU / NonNav IoU : " + f"{global_seg['nav_iou']:.4f} / " + f"{global_seg['non_nav_iou']:.4f}" + ) + print( + f"Unsafe nav global : " + f"{100.0 * global_seg['unsafe_nav_rate']:.4f}%" + ) + print( + f"Status acc / macroF1 : " + f"{global_status['acc']:.4f} / " + f"{global_status['macro_f1']:.4f}" + ) + print( + f"Report : {report_path}" + ) + print( + f"Candidates : {candidates_path}" + ) + print( + f"Summary : {summary_path}" + ) + + if not args.report_only: + print( + f"Revisão física : {review_group_root}" + ) + print( + f"Source manifest : " + f"{out_root / 'review_source_manifest.csv'}" + ) + + print("=" * 94) + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/oak-d/config.json b/Python/OAK/datasets/oak-d/config.json new file mode 100644 index 000000000..ef85631d1 --- /dev/null +++ b/Python/OAK/datasets/oak-d/config.json @@ -0,0 +1,185 @@ +{ + "camera": "oak-d", + "modelo": "segformer_b0", + "model_name": "corredores_v2", + "main_class_name": "navegavel", + "model_to_use": "geral", + "ckpt_test": "best_operational", + + "raw_size": [1920, 1080], + "resolucao": [1024, 576], + "roi_inicio": 0.0, + "roi_tamanho": 1.0, + + "channels": 3, + "backbone": "nvidia/mit-b0", + + "label_classes": [ + "Parado", + "EntrandoRua", + "CaminhandoRua", + "SaindoRua", + "Manobrando", + "Direcionando", + "RetornandoBase", + "Indefinido" + ], + + "corridor_training": { + "data": { + "strict_labels": true, + "strict_pipeline_contract": true, + "preflight_max_samples": 0 + }, + + "augmentation": { + "enabled": true, + "horizontal_flip_p": 0.50, + "affine_p": 0.60, + "rotate_deg": 4.0, + "scale_min": 0.92, + "scale_max": 1.08, + "translate_frac": 0.03, + + "global_gain_p": 0.45, + "global_gain_min": 0.82, + "global_gain_max": 1.18, + + "contrast_p": 0.35, + "contrast_min": 0.85, + "contrast_max": 1.15, + + "gamma_p": 0.30, + "gamma_min": 0.82, + "gamma_max": 1.18, + + "rgb_gain_p": 0.25, + "rgb_gain_min": 0.94, + "rgb_gain_max": 1.06, + + "shadow_p": 0.35, + "shadow_strength_min": 0.12, + "shadow_strength_max": 0.42, + + "noise_p": 0.18, + "noise_sigma_min": 0.003, + "noise_sigma_max": 0.018, + + "blur_p": 0.12, + "blur_kernel": 3, + + "occlusion_p": 0.00, + "occlusion_min_frac": 0.04, + "occlusion_max_frac": 0.15, + "ramp_epochs": 6 + }, + + "sampler": { + "mode": "joint", + "joint_power": 0.30, + "status_power": 0.12, + "group_power": 0.10, + "max_weight_ratio": 5.0, + "samples_per_epoch": 0 + }, + + "class_weighting": { + "seg_method": "log_inverse", + "seg_log_offset": 1.02, + "seg_min_weight": 0.35, + "seg_max_weight": 3.0, + "seg_max_samples": 1200, + + "status_enabled": true, + "status_power": 0.30, + "status_min_weight": 0.50, + "status_max_weight": 3.0 + }, + + "loss": { + "seg_ce_weight": 0.65, + "seg_dice_weight": 0.35, + "boundary_boost": 0.15, + "unsafe_nav_weight": 0.06, + + "status_weight": 0.40, + "status_ramp_epochs": 8, + "status_label_smoothing": 0.03 + }, + + "status_head": { + "pool_h": 3, + "pool_w": 4, + "hidden": 256, + "dropout": 0.20, + "detach_seg_summary": true, + "seg_summary": "probabilities" + }, + + "optimizer": { + "encoder_lr_mult": 0.50, + "decoder_lr_mult": 1.00, + "status_lr_mult": 2.00, + "freeze_encoder_epochs": 2, + "no_decay_bias": true, + "no_decay_norm": true, + "betas": [0.9, 0.999], + "eps": 1e-8 + }, + + "scheduler": { + "mode": "poly", + "warmup_ratio": 0.05, + "warmup_start_factor": 0.10, + "poly_power": 1.0, + "min_lr_ratio": 0.02 + }, + + "optimization": { + "grad_clip_norm": 1.0, + "matmul_precision": "high", + "cudnn_benchmark": true, + "persistent_workers": true, + "prefetch_factor": 2 + }, + + "score": { + "nav_iou": 0.40, + "nav_f1": 0.20, + "status_macro_f1": 0.20, + "safety": 0.20 + }, + + "operational_gate": { + "enabled": true, + "min_epoch": 5, + + "min_nav_iou": 0.70, + "min_nav_f1": 0.80, + "min_non_nav_iou": 0.60, + "min_status_macro_f1": 0.60, + "max_unsafe_nav_rate": 0.03, + + "group_unsafe": { + "enabled": true, + "min_non_nav_pixels": 5000, + "max_rate": 0.05, + "max_rate_by_group": { + "naonavegavel": 0.04, + "navegavel_naonavegavel": 0.05 + } + }, + + "status_recall": { + "enabled": false, + "min_support": 5, + "minimums": {} + } + }, + + "checkpoint": { + "early_stop_patience": 18, + "early_stop_min_delta": 0.0003 + } + } +} diff --git a/Python/OAK/datasets/oak-d/roi_seg_dataset.py b/Python/OAK/datasets/oak-d/roi_seg_dataset.py new file mode 100644 index 000000000..8b9b2221b --- /dev/null +++ b/Python/OAK/datasets/oak-d/roi_seg_dataset.py @@ -0,0 +1,148 @@ +# -*- coding: utf-8 -*- +from PIL import Image +import glob +import os +import cv2 +import numpy as np +import torch +from torch.utils.data import Dataset + +from utils import carregar_labelmap_completo, compute_roi_indices, resize_keep_width + +IMG_EXTS = (".jpg", ".jpeg", ".png") +MSK_EXTS = (".png", ".jpg", ".jpeg") # preferimos .png se existir + +def _is_dir(p): return os.path.isdir(p) +def _is_file(p): return os.path.isfile(p) + +def _list_groups(group_root): + if not _is_dir(group_root): return [] + out = [] + for g in sorted(os.listdir(group_root)): + gdir = os.path.join(group_root, g) + if not _is_dir(gdir): + continue + if _is_dir(os.path.join(gdir, "images")) and _is_dir(os.path.join(gdir, "masks")): + out.append(g) + return out + +def _mask_for_base(msk_dir, base): + """Encontra a máscara que casa com o base, priorizando .png.""" + best = None + for ext in MSK_EXTS: + cand = os.path.join(msk_dir, base + ext) + if _is_file(cand): + if best is None: best = cand + # mantém .png se aparecer depois + if os.path.splitext(cand)[1].lower() == ".png": + return cand + return best + +def _collect_pairs_legacy(root): + """root/{images,masks}""" + img_dir = os.path.join(root, "images") + msk_dir = os.path.join(root, "masks") + imgs = [] + msks = [] + for p in sorted(glob.glob(os.path.join(img_dir, "*"))): + base, ext = os.path.splitext(os.path.basename(p)) + if ext.lower() not in IMG_EXTS: + continue + m = _mask_for_base(msk_dir, base) + if m: + imgs.append(p) + msks.append(m) + return imgs, msks + +def _collect_pairs_grouped(root): + """Suporta: + - root/group//{images,masks} + - root//{images,masks} (quando root já é '.../group') + """ + # case A: root tem subpasta 'group' + group_root = os.path.join(root, "group") + if not _is_dir(group_root): + # case B: root JÁ É a pasta 'group' + group_root = root + + groups = _list_groups(group_root) + imgs, msks = [], [] + for g in groups: + img_dir = os.path.join(group_root, g, "images") + msk_dir = os.path.join(group_root, g, "masks") + for p in sorted(glob.glob(os.path.join(img_dir, "*"))): + base, ext = os.path.splitext(os.path.basename(p)) + if ext.lower() not in IMG_EXTS: + continue + m = _mask_for_base(msk_dir, base) + if m: + imgs.append(p) + msks.append(m) + return imgs, msks + +class ROISegDataset(Dataset): + """ + Compatível com o dataset original, mas agora aceita: + - root = '.../split/train' (com 'group' dentro) + - root = '.../split/train/group' + - root = '.../split/train/' (ainda funciona via legado se tiver images/masks) + - root legado = '.../split/train' com 'images' e 'masks' diretamente + """ + def __init__(self, root, out_dir, zona_inicio, faixa_atuacao, input_w=384, min_input_h=96, labelmap_path="labelmap.txt", + mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)): + + self.out_dir = out_dir + self.zona_inicio = zona_inicio + self.faixa_atuacao = faixa_atuacao + self.input_w = input_w + self.min_input_h = min_input_h + self.mean = np.array(mean, dtype=np.float32).reshape(1, 1, 3) + self.std = np.array(std, dtype=np.float32).reshape(1, 1, 3) + + # tenta agrupar; se não encontrar, cai pro legado + imgs, msks = _collect_pairs_grouped(root) + if not imgs: + imgs, msks = _collect_pairs_legacy(root) + + assert len(imgs) == len(msks) and len(imgs) > 0, f"Nenhuma imagem/máscara encontrada em {root}" + self.img_paths = imgs + self.msk_paths = msks + + # labelmap + _, _, self.classes, self.ignore_rgb = carregar_labelmap_completo(labelmap_path) + # tipicamente ignore_rgb é (255,255,255) + self.ignore_id = int(self.ignore_rgb[0]) if isinstance(self.ignore_rgb, (list, tuple)) else int(self.ignore_rgb) + + def __len__(self): + return len(self.img_paths) + + def __getitem__(self, idx): + img_rgb = np.array(Image.open(self.img_paths[idx]).convert("RGB")) + # Máscara como escala de cinza (IDs já foram normalizados na etapa de normalize) + msk_grayscale = np.array(Image.open(self.msk_paths[idx]).convert("L")) + + H, W = img_rgb.shape[:2] + y_fim, y_inicio = compute_roi_indices(H, self.zona_inicio, self.faixa_atuacao) + + img_roi = img_rgb[y_fim:y_inicio, 0:W] + msk_roi = msk_grayscale[y_fim:y_inicio, 0:W] + + img_in = resize_keep_width(img_roi, self.input_w, self.min_input_h, cv2.INTER_AREA) + msk_ids = resize_keep_width(msk_roi, self.input_w, self.min_input_h, cv2.INTER_NEAREST) + + valores_validos = list(range(len(self.classes))) + [255] + msk_ids[np.isin(msk_ids, valores_validos, invert=True)] = self.ignore_id + + if np.all(msk_ids == self.ignore_id): + raise ValueError(f"Máscara {self.msk_paths[idx]} está só com valor de ignore ({self.ignore_id})") + + img_f = img_in.astype(np.float32) / 255.0 + img_f = (img_f - self.mean) / self.std + img_chw = np.transpose(img_f, (2, 0, 1)) + + if idx == 0: + os.makedirs(self.out_dir, exist_ok=True) + cv2.imwrite(os.path.join(self.out_dir, "debug_roi_input.png"), cv2.cvtColor(img_roi, cv2.COLOR_RGB2BGR)) + cv2.imwrite(os.path.join(self.out_dir, "debug_roi_mask.png"), msk_roi) + + return torch.from_numpy(img_chw).float(), torch.from_numpy(msk_ids.astype(np.int64)) diff --git a/Python/OAK/datasets/oak-d/utils.py b/Python/OAK/datasets/oak-d/utils.py new file mode 100644 index 000000000..9207e7c27 --- /dev/null +++ b/Python/OAK/datasets/oak-d/utils.py @@ -0,0 +1,144 @@ +import cv2 +import numpy as np + +# ---------------------------- +# Helpers LABELMAP +# ---------------------------- +def carregar_labelmap_completo(caminho): + cor_para_id = {} + id_para_nome = {} + cores_rgb = [] + + with open(caminho, 'r') as arquivo: + idx = 0 + for linha in arquivo: + if linha.startswith("#") or not linha.strip(): + continue + partes = linha.strip().split(':') + if len(partes) >= 2: + nome_classe, cor_rgb_str = partes[0], partes[1] + r, g, b = map(int, cor_rgb_str.split(',')) + cor_rgb = (r, g, b) + + if nome_classe.lower() == "ignore": + ignore_rgb = cor_rgb + continue # NÃO adiciona ignore no LUT de classes + + cor_para_id[cor_rgb] = idx + cores_rgb.append(cor_rgb) + id_para_nome[idx] = nome_classe + idx += 1 + + #print(f"Mapa: {cor_para_id}") + #print(f"Colormap RGB: {cores_rgb}") + #print(f"Classes: {id_para_nome}") + #print(f"Ignore RGB: {ignore_rgb}") + + return cor_para_id, cores_rgb, id_para_nome, ignore_rgb + +def converter_mask_rgb_para_ids(img_rgb, mapa_rgb, ignore_id): + # Cria um mapa 256^3 para IDs (usa int32 para indexar) + lut = np.full((256**3,), ignore_id, dtype=np.uint8) + + for cor, classe_id in mapa_rgb.items(): + r, g, b = cor + lut[(r << 16) + (g << 8) + b] = classe_id + + # Converte RGB para índice único + flat_idx = (img_rgb[:,:,0].astype(np.int32) << 16) + \ + (img_rgb[:,:,1].astype(np.int32) << 8) + \ + img_rgb[:,:,2].astype(np.int32) + + # Aplica LUT vetorizada + return lut[flat_idx] + +def converter_mask_ids_para_rgb(mask_ids: np.ndarray, colormap_rgb: list, ignore_id: int = 255) -> np.ndarray: + # Criar lookup table (256 cores possíveis) + lut = np.zeros((256, 3), dtype=np.uint8) + for i, color in enumerate(colormap_rgb): + lut[i] = color + lut[ignore_id] = (255, 255, 255) + + # Aplicar LUT direto (vetorizado) + return lut[mask_ids] + +def converter_mask_ids_para_bgr(mask_ids: np.ndarray, colormap_rgb: list, ignore_id: int = 255) -> np.ndarray: + """ + Converte máscara de IDs para imagem BGR (uint8), + pronta para uso com OpenCV. + """ + lut = np.zeros((256, 3), dtype=np.uint8) + + for i, (r, g, b) in enumerate(colormap_rgb): + lut[i] = (b, g, r) # RGB -> BGR + + lut[ignore_id] = (255, 255, 255) # branco em BGR = RGB + + return lut[mask_ids] + +def desenhar_legenda_vertical(colormap_rgb, classes, largura=200): + """ + Retorna uma imagem com a legenda das classes (cor + nome) + """ + nomes_classes = [classes[i] for i in range(len(classes))] + altura_por_classe = 30 + altura_total = altura_por_classe * len(colormap_rgb) + legenda = np.ones((altura_total, largura, 3), dtype=np.uint8) * 255 + + for idx, (rgb, nome) in enumerate(zip(colormap_rgb, nomes_classes)): + y = idx * altura_por_classe + color = tuple(int(c) for c in rgb) + cv2.rectangle(legenda, (10, y + 5), (30, y + 25), color, -1) + cv2.putText(legenda, nome, (40, y + 20), + cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 0, 0), 1, cv2.LINE_AA) + + return legenda + +def desenhar_legenda_horizontal(colormap_rgb, classes, altura=30, largura_por_classe=120): + """ + Retorna uma imagem com a legenda das classes (cor + nome), em uma única linha horizontal + """ + nomes_classes = [classes[i] for i in range(len(classes))] + largura_total = largura_por_classe * len(colormap_rgb) + legenda = np.ones((altura, largura_total, 3), dtype=np.uint8) * 255 # faixa branca + + for idx, (rgb, nome) in enumerate(zip(colormap_rgb, nomes_classes)): + x = idx * largura_por_classe + color = tuple(int(c) for c in rgb) + # Retângulo colorido + cv2.rectangle(legenda, (x + 10, 5), (x + 30, 25), color, -1) + # Texto da classe + cv2.putText(legenda, nome, (x + 35, 20), + cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 0, 0), 1, cv2.LINE_AA) + + return legenda + +def _infer_ignore_id(ignore_rgb, default_id=255): + import numpy as _np + if isinstance(ignore_rgb, (list, tuple)): + if len(ignore_rgb) == 1 and isinstance(ignore_rgb[0], (int, _np.integer)): + return int(ignore_rgb[0]) + if len(ignore_rgb) == 3: + return default_id + if isinstance(ignore_rgb, (int, _np.integer)): + return int(ignore_rgb) + return default_id + +# ---------------------------- +# Helpers ROI +# ---------------------------- +def compute_roi_indices(H: int, zona_inicio: float, faixa_atuacao: float): + y_inicio = int((1.0 - zona_inicio) * H) + y_fim = int((1.0 - (zona_inicio + faixa_atuacao)) * H) + y_fim = max(0, min(H, y_fim)) + y_inicio = max(0, min(H, y_inicio)) + if y_fim >= y_inicio: + y_fim = max(0, y_inicio - 1) + return y_fim, y_inicio + +def resize_keep_width(img: np.ndarray, new_w: int, min_h: int, interpolation: int) -> np.ndarray: + h, w = img.shape[:2] + new_h = int(round(new_w * (h / w))) + if min_h is not None and new_h < min_h: + new_h = min_h + return cv2.resize(img, (new_w, new_h), interpolation=interpolation)