CoolFace
Apppublic

Allex21/Lora-trainer-all

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
lora_training.py345 linesDownload Raw Back to root
1import os2import uuid3import json4import threading5import subprocess6import time7from datetime import datetime8from flask import Blueprint, request, jsonify, send_file9from werkzeug.utils import secure_filename10import shutil11from src.lora_trainer import create_lora_trainer, validate_training_config12 13lora_bp = Blueprint('lora', __name__)14 15# Configurações16UPLOAD_FOLDER = '/tmp/lora_uploads'17TRAINING_FOLDER = '/tmp/lora_training'18RESULTS_FOLDER = '/tmp/lora_results'19ALLOWED_EXTENSIONS = {'png', 'jpg', 'jpeg', 'webp', 'bmp'}20 21# Armazenamento em memória para status de treinamento22training_status = {}23 24def allowed_file(filename):25    return '.' in filename and filename.rsplit('.', 1)[1].lower() in ALLOWED_EXTENSIONS26 27def ensure_directories():28    """Garante que os diretórios necessários existam"""29    for folder in [UPLOAD_FOLDER, TRAINING_FOLDER, RESULTS_FOLDER]:30        os.makedirs(folder, exist_ok=True)31 32def run_training_process(training_id, config):33    """Executa o processo de treinamento LoRA em thread separada"""34    try:35        # Atualizar status36        training_status[training_id]['status'] = 'preparing'37        training_status[training_id]['progress'] = 1038        training_status[training_id]['message'] = 'Preparando ambiente de treinamento...'39        40        # Criar diretório de treinamento41        training_dir = os.path.join(TRAINING_FOLDER, training_id)42        os.makedirs(training_dir, exist_ok=True)43        44        # Copiar imagens para diretório de treinamento45        images_dir = os.path.join(training_dir, 'images')46        os.makedirs(images_dir, exist_ok=True)47        48        upload_dir = os.path.join(UPLOAD_FOLDER, training_id)49        for filename in os.listdir(upload_dir):50            if allowed_file(filename):51                src = os.path.join(upload_dir, filename)52                dst = os.path.join(images_dir, filename)53                shutil.copy2(src, dst)54        55        training_status[training_id]['progress'] = 2056        training_status[training_id]['message'] = 'Preparando dataset...'57        58        # Configurar paths para treinamento59        output_dir = os.path.join(RESULTS_FOLDER, training_id)60        os.makedirs(output_dir, exist_ok=True)61        62        config['images_dir'] = images_dir63        config['output_dir'] = output_dir64        65        # Validar configuração66        is_valid, validation_message = validate_training_config(config)67        if not is_valid:68            raise Exception(f"Configuração inválida: {validation_message}")69        70        training_status[training_id]['progress'] = 3071        training_status[training_id]['message'] = 'Iniciando treinamento LoRA...'72        73        # Callback para atualizar progresso74        def progress_callback(progress, message):75            training_status[training_id]['progress'] = max(30, min(90, progress))76            training_status[training_id]['message'] = message77            training_status[training_id]['logs'] = training_status[training_id].get('logs', []) + [message]78        79        # Criar e executar trainer80        trainer = create_lora_trainer(config)81        82        # Executar treinamento83        trainer.train(progress_callback=progress_callback)84        85        # Atualizar logs86        training_status[training_id]['logs'] = trainer.training_logs87        88        training_status[training_id]['status'] = 'completed'89        training_status[training_id]['progress'] = 10090        training_status[training_id]['message'] = 'Treinamento concluído com sucesso!'91        training_status[training_id]['completed'] = True92        93        # Criar arquivos adicionais94        create_additional_files(training_id, config, output_dir)95        96        # Criar links de download97        download_links = []98        for filename in os.listdir(output_dir):99            if os.path.isfile(os.path.join(output_dir, filename)):100                download_links.append({101                    'name': filename,102                    'url': f'/api/download/{training_id}/{filename}'103                })104        105        training_status[training_id]['download_links'] = download_links106        training_status[training_id]['trigger_word'] = config['trigger_word']107        108    except Exception as e:109        training_status[training_id]['status'] = 'error'110        training_status[training_id]['error'] = str(e)111        training_status[training_id]['message'] = f'Erro durante treinamento: {str(e)}'112 113def create_additional_files(training_id, config, output_dir):114    """Cria arquivos adicionais de resultado"""115    116    # Criar README com instruções detalhadas117    readme_content = f'''# LoRA: {config['character_name']}118 119## Informações do Treinamento120- **Personagem**: {config['character_name']}121- **Trigger Word**: {config['trigger_word']}122- **Resolução**: {config['resolution']}x{config['resolution']}123- **Rank**: {config['rank']}124- **Learning Rate**: {config['learning_rate']}125- **Épocas**: {config['epochs']}126- **Data de Treinamento**: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}127 128## Como Usar129 130### ComfyUI1311. Coloque o arquivo `pytorch_lora_weights.safetensors` na pasta `models/loras/` do ComfyUI1322. No workflow, adicione um nó "Load LoRA"1333. Selecione o arquivo LoRA1344. Use a trigger word "{config['trigger_word']}" em seus prompts1355. Ajuste o peso entre 0.7 e 1.0136 137### Automatic11111381. Coloque o arquivo `pytorch_lora_weights.safetensors` na pasta `models/Lora/`1392. Use a sintaxe `<lora:pytorch_lora_weights:0.8>` no prompt1403. Inclua a trigger word "{config['trigger_word']}" no prompt141 142### SeaArt1431. Faça upload do arquivo LoRA na plataforma1442. Selecione o LoRA em suas gerações1453. Use a trigger word "{config['trigger_word']}" no prompt146 147## Exemplos de Prompts148- "{config['trigger_word']}, portrait, high quality"149- "{config['trigger_word']}, full body, standing, detailed"150- "{config['trigger_word']}, close-up, beautiful lighting"151- "{config['trigger_word']}, anime style, colorful"152- "{config['trigger_word']}, realistic, photographic"153 154## Dicas de Uso155- Use peso entre 0.7-1.0 para melhores resultados156- Combine com outros LoRAs para estilos específicos157- Experimente diferentes CFG scales (7-12)158- Para resultados mais consistentes, use a trigger word no início do prompt159- Ajuste o peso do LoRA conforme necessário para cada geração160 161## Compatibilidade162Este LoRA é compatível com:163- Stable Diffusion 1.5164- ComfyUI165- Automatic1111 WebUI166- SeaArt167- Fooocus168- InvokeAI169- E outras ferramentas que suportam LoRA170 171## Suporte172Para dúvidas ou problemas, consulte a documentação da ferramenta utilizada.173'''174    175    readme_file = os.path.join(output_dir, 'README.md')176    with open(readme_file, 'w', encoding='utf-8') as f:177        f.write(readme_content)178    179    # Criar arquivo de metadados180    metadata = {181        "model_name": config['character_name'],182        "trigger_word": config['trigger_word'],183        "base_model": "runwayml/stable-diffusion-v1-5",184        "training_config": config,185        "created_at": datetime.now().isoformat(),186        "version": "1.0",187        "type": "character_lora",188        "tags": ["character", "lora", "consistent", config['character_name']],189        "description": f"LoRA treinado para o personagem {config['character_name']} usando a trigger word '{config['trigger_word']}'",190        "usage_instructions": {191            "trigger_word": config['trigger_word'],192            "recommended_weight": "0.7-1.0",193            "compatible_models": ["SD1.5"],194            "example_prompts": [195                f"{config['trigger_word']}, portrait, high quality",196                f"{config['trigger_word']}, full body, detailed",197                f"{config['trigger_word']}, close-up, beautiful lighting"198            ]199        }200    }201    202    metadata_file = os.path.join(output_dir, 'metadata.json')203    with open(metadata_file, 'w', encoding='utf-8') as f:204        json.dump(metadata, f, indent=2, ensure_ascii=False)205 206@lora_bp.route('/train', methods=['POST'])207def start_training():208    """Inicia o treinamento LoRA"""209    try:210        ensure_directories()211        212        # Gerar ID único para o treinamento213        training_id = str(uuid.uuid4())214        215        # Verificar se há imagens216        if 'images' not in request.files:217            return jsonify({'success': False, 'message': 'Nenhuma imagem foi enviada'}), 400218        219        files = request.files.getlist('images')220        if len(files) < 5:221            return jsonify({'success': False, 'message': 'Mínimo de 5 imagens necessárias'}), 400222        223        # Criar diretório para upload224        upload_dir = os.path.join(UPLOAD_FOLDER, training_id)225        os.makedirs(upload_dir, exist_ok=True)226        227        # Salvar imagens228        saved_files = []229        for file in files:230            if file and allowed_file(file.filename):231                filename = secure_filename(file.filename)232                filepath = os.path.join(upload_dir, filename)233                file.save(filepath)234                saved_files.append(filename)235        236        if len(saved_files) < 5:237            return jsonify({'success': False, 'message': 'Pelo menos 5 imagens válidas são necessárias'}), 400238        239        # Obter configurações240        config = {241            'character_name': request.form.get('character_name', '').strip(),242            'trigger_word': request.form.get('trigger_word', '').strip(),243            'resolution': request.form.get('resolution', '512'),244            'learning_rate': request.form.get('learning_rate', '1e-4'),245            'rank': request.form.get('rank', '16'),246            'epochs': request.form.get('epochs', '20'),247            'description': request.form.get('description', '').strip(),248            'images': saved_files,249            'training_id': training_id,250            'created_at': datetime.now().isoformat()251        }252        253        # Validar configurações obrigatórias254        if not config['character_name']:255            return jsonify({'success': False, 'message': 'Nome do personagem é obrigatório'}), 400256        257        if not config['trigger_word']:258            return jsonify({'success': False, 'message': 'Trigger word é obrigatória'}), 400259        260        # Inicializar status do treinamento261        training_status[training_id] = {262            'status': 'starting',263            'progress': 0,264            'message': 'Iniciando treinamento...',265            'logs': [],266            'completed': False,267            'error': None,268            'config': config269        }270        271        # Iniciar treinamento em thread separada272        training_thread = threading.Thread(273            target=run_training_process,274            args=(training_id, config)275        )276        training_thread.daemon = True277        training_thread.start()278        279        return jsonify({280            'success': True,281            'training_id': training_id,282            'message': 'Treinamento iniciado com sucesso'283        })284        285    except Exception as e:286        return jsonify({'success': False, 'message': f'Erro interno: {str(e)}'}), 500287 288@lora_bp.route('/training-status/<training_id>', methods=['GET'])289def get_training_status(training_id):290    """Retorna o status do treinamento"""291    if training_id not in training_status:292        return jsonify({'error': 'Treinamento não encontrado'}), 404293    294    return jsonify(training_status[training_id])295 296@lora_bp.route('/download/<training_id>/<filename>', methods=['GET'])297def download_file(training_id, filename):298    """Download de arquivos de resultado"""299    try:300        result_dir = os.path.join(RESULTS_FOLDER, training_id)301        file_path = os.path.join(result_dir, filename)302        303        if not os.path.exists(file_path):304            return jsonify({'error': 'Arquivo não encontrado'}), 404305        306        return send_file(file_path, as_attachment=True, download_name=filename)307        308    except Exception as e:309        return jsonify({'error': f'Erro ao baixar arquivo: {str(e)}'}), 500310 311@lora_bp.route('/trainings', methods=['GET'])312def list_trainings():313    """Lista todos os treinamentos"""314    trainings = []315    for training_id, status in training_status.items():316        trainings.append({317            'id': training_id,318            'character_name': status.get('config', {}).get('character_name', 'Desconhecido'),319            'status': status.get('status', 'unknown'),320            'progress': status.get('progress', 0),321            'created_at': status.get('config', {}).get('created_at', '')322        })323    324    return jsonify({'trainings': trainings})325 326@lora_bp.route('/delete-training/<training_id>', methods=['DELETE'])327def delete_training(training_id):328    """Remove um treinamento e seus arquivos"""329    try:330        # Remover do status331        if training_id in training_status:332            del training_status[training_id]333        334        # Remover diretórios335        for base_dir in [UPLOAD_FOLDER, TRAINING_FOLDER, RESULTS_FOLDER]:336            training_dir = os.path.join(base_dir, training_id)337            if os.path.exists(training_dir):338                shutil.rmtree(training_dir)339        340        return jsonify({'success': True, 'message': 'Treinamento removido com sucesso'})341        342    except Exception as e:343        return jsonify({'success': False, 'message': f'Erro ao remover treinamento: {str(e)}'}), 500344 345