tlc-agent-scientifique / optimizer.py
Gilbra's picture
chnage provider
ed78f98
Raw
History Blame Contribute Delete
2.99 kB
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)
]