agrobot_base/Treinamento/models/ruas/process_Ruas_dynamic.py

224 lines
8.5 KiB
Python

import torch
from PIL import Image
import torchvision.transforms as T
import numpy as np
import cv2
import json
import time
from flask import Flask, Response, request
import threading
import argparse
import socket
import os
# Argumentos do script
parser = argparse.ArgumentParser(description='Transmite vídeo processado por modelo de deep learning.')
parser.add_argument('--max_readings', type=int, help='Número máximo de leituras a serem mantidas.', required=True)
parser.add_argument('--port', type=int, help='Porta para o servidor Flask.', required=True)
parser.add_argument('--url', type=str, help='URL para o stream de vídeo.', required=True)
parser.add_argument('--output_file', type=str, help='Arquivo de saída para salvar as detecções.', required=True)
parser.add_argument('--show_lines', type=int, choices=[0, 1], help='Se deve mostrar linhas de detecção.', required=True)
parser.add_argument('--socket_port', type=int, help='Porta para o servidor de socket.', required=True)
parser.add_argument('--camera_index', type=int, help='Índice da câmera a ser usada.', default=0)
args = parser.parse_args()
json_reading_data = None
app = Flask(__name__)
# Supondo que você tenha a estrutura do repositório e o módulo `network` conforme descrito no README
from network.modeling import deeplabv3plus_resnet50 as deeplabv3_model
# Configurações Iniciais
NUM_CLASSES = 4 # Pascal VOC possui 3 classes + 1 para o fundo
OUTPUT_STRIDE = 16 # Valor comum para DeepLab
MODEL_PATH = 'backup/ruasModel_final.pth' # Caminho para o modelo pré-treinado
# Função para carregar o modelo
def load_model(model_path):
model = deeplabv3_model(num_classes=NUM_CLASSES, output_stride=OUTPUT_STRIDE)
model.load_state_dict(torch.load(model_path), strict=False)
model.eval() # Modo de avaliação
return model
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Carregar o modelo
model = load_model(MODEL_PATH)
model.to(device)
# Função modificada para processar um frame da câmera
def segment_frame(frame):
# Converte o frame do OpenCV (BGR) para o formato RGB
image = Image.fromarray(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB))
transform = T.Compose([
T.Resize(520),
T.ToTensor(),
T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
input_tensor = transform(image).unsqueeze(0).to(device)
with torch.no_grad():
output = model(input_tensor)
output_predictions = output.max(1)[1].squeeze().detach().cpu().numpy()
generate_and_save_json(output_predictions, args.output_file)
return output_predictions
def read_labelmap(path):
label_colors = []
class_names = []
with open(path, 'r') as file:
for line in file.readlines():
# Ignora linhas comentadas
if line.startswith('#'):
continue
parts = line.strip().split(':')
if len(parts) >= 2:
label = parts[0].strip()
color = tuple(map(int, parts[1].split(',')))
class_names.append(label)
label_colors.append(color)
return np.array(label_colors), class_names
def generate_and_save_json(image, output_path='contours_data.json', labelmap_path="dataset/labelmap.txt"):
label_colors, class_names = read_labelmap(labelmap_path)
detected_classes = set(np.unique(image))
global json_reading_data
contours_info = []
# Carregar os dados existentes se o arquivo já existir
if os.path.exists(output_path):
with open(output_path, 'r') as f:
try:
existing_data = json.load(f)
if type(existing_data) is list:
contours_info.extend(existing_data)
except json.JSONDecodeError:
print("Erro ao decodificar o JSON existente. Um novo arquivo será criado.")
timestamp = time.time()
json_data = {'timestamp': timestamp, 'Classes': []}
for l in detected_classes:
if l < len(label_colors): # Verifica se o índice está dentro do intervalo das cores definidas
json_data['Classes'].append({'Classe': class_names[l], 'Contornos': []})
# Extrai a máscara para a classe atual
mask = (image == l).astype(np.uint8) * 255
contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
for contour in contours:
# Verifica se o contorno não está vazio e tem a dimensão adequada
if contour.size > 0 and len(contour.shape) == 3 and contour.shape[1] == 2:
# Calcula os pontos extremos e intermediários para o contorno atual
json_data['Contornos'].append(contour.squeeze().tolist())
contours_info.append(json_data)
# Limita a quantidade de registros a serem salvos com base em max_readings
contours_info = contours_info[-args.max_readings:]
# Salva os dados em um arquivo JSON
with open(output_path, 'w') as f:
json.dump(contours_info, f, indent=4)
json_reading_data = contours_info
# Função para decodificar e aplicar o mapa de segmentação em um frame
def apply_segmentation_overlay(frame, output_predictions, labelmap_path="dataset/labelmap.txt"):
label_colors, class_names = read_labelmap(labelmap_path)
nc = len(label_colors)
height, width, _ = frame.shape
overlay = np.zeros((height, width, 3), dtype=np.uint8)
# Redimensiona as previsões do modelo para corresponder ao tamanho do frame
output_predictions_resized = cv2.resize(output_predictions, (width, height), interpolation=cv2.INTER_NEAREST)
for l in np.unique(output_predictions_resized):
if l < nc:
mask = output_predictions_resized == l
overlay[mask] = label_colors[l]
# Combinação do frame original com o overlay da segmentação
overlayed_frame = cv2.addWeighted(frame, 0.6, overlay, 0.4, 0)
return overlayed_frame
# Função para processar e transmitir o vídeo
def detect_and_stream(camera_index):
cap = cv2.VideoCapture(camera_index)
while True:
ret, frame = cap.read()
if not ret:
break
# Processa o frame com seu modelo
output_predictions = segment_frame(frame)
# Aplica a segmentação sobre o frame capturado
overlayed_frame = apply_segmentation_overlay(frame, output_predictions)
frame_saida = overlayed_frame if args.show_lines == 1 else frame
# Codifica o frame para JPEG e transmite
ret, buffer = cv2.imencode('.jpg', frame_saida)
frame = buffer.tobytes()
yield (b'--frame\r\n'
b'Content-Type: image/jpeg\r\n\r\n' + frame + b'\r\n')
# Função para lidar com a conexão de cada cliente
def handle_client_connection(client_socket):
try:
while True:
# Enviar o dicionário serializado em JSON
client_socket.send(json.dumps(json_reading_data).encode('utf-8'))
# Aguardar um pouco antes de enviar os próximos dados
time.sleep(0.2)
except socket.error:
print(f"Cliente desconectado.")
finally:
# Fechar a conexão do socket ao sair do loop
client_socket.close()
# Configuração inicial do servidor de socket
def start_server(address='localhost', port=args.socket_port):
server = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
server.bind((address, port))
server.listen()
print(f"Servidor iniciado. Aguardando conexões em {address}:{port}...")
try:
while True:
client_sock, address = server.accept()
print(f"Aceitando conexão de {address[0]}:{address[1]}")
client_handler = threading.Thread(
target=handle_client_connection,
args=(client_sock,)
)
client_handler.start()
finally:
server.close()
# Esta função é para iniciar o servidor de socket em uma thread separada
def run_socket_server():
start_server()
@app.route('/' + args.url, methods=['GET'])
def stream():
return Response(detect_and_stream(args.camera_index),
mimetype='multipart/x-mixed-replace; boundary=frame')
def run_flask_server():
app.run(host='0.0.0.0', port=args.port, threaded=True, debug=False)
# Iniciar o servidor
if __name__ == '__main__':
# Inicia o servidor de socket em uma thread separada
socket_server_thread = threading.Thread(target=run_socket_server)
socket_server_thread.start()
# Inicia o servidor Flask na thread principal
run_flask_server()