from typing import Dict, Any, List import json from utils import ( call_llm, safe_json_loads, normalize_ir, validate_ir ) # ===================================================== # FALLBACK STRATEGY # ===================================================== def fallback_optimization( ir: Dict[str, Any], idx: int ): return { "explanation": f"Fallback optimisation strategy {idx}", "optimized_ir": normalize_ir(ir) } # ===================================================== # OPTIMIZATION # ===================================================== def optimize_ir( ir: Dict[str, Any], num_strategies: int = 3, provider: str = "openai" ) -> List[Dict[str, Any]]: if not validate_ir(ir): return [ fallback_optimization(ir, i + 1) for i in range(num_strategies) ] ir_text = json.dumps( ir, indent=2, ensure_ascii=False ) prompt = f""" Tu es un expert en optimisation de graphes mathématiques. Analyse cette IR. Produis EXACTEMENT {num_strategies} stratégies d'optimisation différentes. Les optimisations possibles incluent : - fusion d'opérations - suppression de redondances - réduction mémoire - vectorisation - parallélisation - approximation numérique - simplification algébrique - réordonnancement IMPORTANT : - retourne UNIQUEMENT du JSON valide - aucune explication hors JSON - aucune balise markdown FORMAT STRICT : [ {{ "explanation": "description courte", "optimized_ir": {{ "name": "Optimized IR", "strategy": "fusion", "nodes": [], "edges": [] }} }} ] IR : {ir_text} """ response = call_llm( prompt, provider=provider, max_tokens=3500 ) parsed = safe_json_loads(response) # ------------------------------------------------- # VALIDATION # ------------------------------------------------- valid_results = [] if isinstance(parsed, list): for item in parsed: if not isinstance(item, dict): continue explanation = item.get( "explanation", "Optimisation automatique" ) optimized_ir = item.get( "optimized_ir", ir ) optimized_ir = normalize_ir( optimized_ir ) valid_results.append({ "explanation": explanation, "optimized_ir": optimized_ir }) # ------------------------------------------------- # SUCCESS # ------------------------------------------------- if len(valid_results) > 0: return valid_results[:num_strategies] # ------------------------------------------------- # FALLBACK # ------------------------------------------------- return [ fallback_optimization(ir, i + 1) for i in range(num_strategies) ]