CoolFace
Apppublic

Yalpha/AGRITECH-META

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
inference.py172 linesDownload Raw Back to root
1"""2inference.py – AgriDecisionEnv v33Strict OpenEnv log format — every line guaranteed.4 5FORMAT:6[START] task=<task> env=<env> model=<model>7[STEP] step=<n> action=<action> reward=<0.00> done=<true|false> error=<null|msg>8[END] success=<true|false> steps=<n> score=<score> rewards=<r1,r2,...>9"""10import os, sys, json, re11from dotenv import load_dotenv12 13load_dotenv()14 15API_BASE_URL       = os.getenv("API_BASE_URL",    "https://router.huggingface.co/v1")16MODEL_NAME         = os.getenv("MODEL_NAME",      "Qwen/Qwen2.5-7B-Instruct")17HF_TOKEN           = os.getenv("HF_TOKEN")18 19# Optional - if you use from_docker_image():20LOCAL_IMAGE_NAME   = os.getenv("LOCAL_IMAGE_NAME")21TASK               = os.environ.get("AGRI_TASK",       "hard")22SCENARIO           = os.environ.get("AGRI_SCENARIO",   "default")23USE_HARDCODED_PLAN = os.environ.get("USE_HARDCODED_PLAN", "false").lower() == "true"24MAX_STEPS          = {"easy": 1, "medium": 3, "hard": 5}.get(TASK, 5)25ENV_NAME           = "AgriDecisionEnv-v3"26 27sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))28from env import AgriEnv29from models import Action30 31# Pre-computed optimal plan for default scenario, hard task.32# Weather sequence is deterministic: normal, normal, rainy, drought, normal.33HARDCODED_PLAN = [34    Action(crop="wheat", fertilizer=0.20, irrigation=0.10),35    Action(crop="none",  fertilizer=0.20, irrigation=0.05),36    Action(crop="wheat", fertilizer=0.20, irrigation=0.05),37    Action(crop="none",  fertilizer=0.20, irrigation=0.40),38    Action(crop="wheat", fertilizer=0.20, irrigation=0.24),39]40 41try:42    from openai import OpenAI43    _api_key = HF_TOKEN or "placeholder"44    _client  = OpenAI(base_url=API_BASE_URL, api_key=_api_key)45except ImportError:46    _client = None47 48 49def _build_prompt(obs: dict) -> str:50    return (51        f"Given the farm state:\n"52        f"nitrogen: {obs['nitrogen']}\n"53        f"moisture: {obs['moisture']}\n"54        f"soil_quality: {obs['soil_quality']}\n"55        f"last_crop: {obs['last_crop']}\n"56        f"season: {obs['season']}\n"57        f"weather: {obs['weather']}\n"58        f"groundwater: {obs['groundwater']}\n"59        f"budget: {obs['budget']}\n\n"60        f"Rules:\n"61        f"- Do NOT repeat last_crop (monocrop penalty -0.12).\n"62        f"- If nitrogen < 0.30 use crop=none to recover soil.\n"63        f"- In drought weather prefer wheat and keep irrigation low.\n"64        f"- In rainy weather reduce irrigation (rain provides +0.12 moisture).\n"65        f"- Budget costs: 20 fixed + fertilizer*15 + irrigation*10 per step.\n"66        f"- Keep fertilizer <= 0.55 and irrigation <= 0.60 to avoid penalties.\n"67        f"- You earn a +0.15 bonus if nitrogen >= 0.45 AND moisture >= 0.35.\n\n"68        f"Suggest the best action. Return ONLY in this exact format:\n"69        f"crop: <rice|wheat|none>\n"70        f"fertilizer: <float 0.0-1.0>\n"71        f"irrigation: <float 0.0-1.0>"72    )73 74 75def _parse_response(text: str) -> Action:76    crop       = re.search(r"crop\s*:\s*(rice|wheat|none)", text, re.I)77    fertilizer = re.search(r"fertilizer\s*:\s*([0-9]*\.?[0-9]+)", text, re.I)78    irrigation = re.search(r"irrigation\s*:\s*([0-9]*\.?[0-9]+)", text, re.I)79    return Action(80        crop       = crop.group(1).lower() if crop else "wheat",81        fertilizer = float(fertilizer.group(1)) if fertilizer else 0.20,82        irrigation = float(irrigation.group(1)) if irrigation else 0.10,83    )84 85 86def _llm_action(obs_dict: dict):87    """Returns (Action, error_str|None)."""88    if _client is None:89        return None, "openai not installed"90    try:91        resp = _client.chat.completions.create(92            model=MODEL_NAME,93            messages=[94                {"role": "system", "content": "You are a sustainable farming AI agent. Follow all rules exactly."},95                {"role": "user",   "content": _build_prompt(obs_dict)},96            ],97            max_tokens=64,98            temperature=0.0,99        )100        raw = resp.choices[0].message.content.strip()101        return _parse_response(raw), None102    except Exception as e:103        return None, str(e)104 105 106def _fallback_action(obs_dict: dict) -> Action:107    from baseline_agents import rule_based_policy108    from models import Observation109    return rule_based_policy(Observation(**obs_dict))110 111 112def run_inference():113    env     = AgriEnv(scenario=SCENARIO, seed=42)114    obs     = env.reset()115    rewards = []116    success = True117 118    print(f"[START] task={TASK} env={ENV_NAME} model={MODEL_NAME}")119 120    for step in range(MAX_STEPS):121        obs_dict  = obs.model_dump() if hasattr(obs, "model_dump") else vars(obs)122        error_msg = "null"123 124        if USE_HARDCODED_PLAN and step < len(HARDCODED_PLAN):125            action = HARDCODED_PLAN[step]126        else:127            action, err = _llm_action(obs_dict)128            if action is None:129                action    = _fallback_action(obs_dict)130                error_msg = err or "fallback"131                if err:132                    success = False133 134        try:135            obs, reward, done, _ = env.step(action)136            rewards.append(reward)137            print(138                f"[STEP] step={step + 1}"139                f" action={json.dumps(action.model_dump() if hasattr(action, 'model_dump') else vars(action))}"140                f" reward={reward:.2f}"141                f" done={str(done).lower()}"142                f" error={error_msg}"143            )144        except Exception as e:145            success = False146            print(147                f"[STEP] step={step + 1}"148                f" action=null"149                f" reward=0.00"150                f" done=true"151                f" error={e}"152            )153            done = True154 155        if done:156            break157 158    score = round(sum(rewards) / max(len(rewards), 1), 4) if rewards else 0.0159 160    rewards_str = ",".join(f"{r:.2f}" for r in rewards)161    print(162        f"[END] success={str(success).lower()}"163        f" steps={len(rewards)}"164        f" score={score}"165        f" rewards={rewards_str}"166    )167    return score168 169 170if __name__ == "__main__":171    run_inference()172