From c51b2c790057fa96bda490c1fb59e83219a121d2 Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 8 Nov 2025 19:23:29 +0000 Subject: [PATCH] Update guarded_ai.py to support classifier_model and feedback_model Changes: - Enhanced get_openai_client_and_model() to support MODEL_X references - Added model parameter (default "MODEL_1") to all AI functions: - categorize_response() - generate_ai_feedback() - provide_feedback() - provide_feedback_prompts() - translate_text() - Updated simulate_activity() to: - Read classifier_model and feedback_model from YAML - Support step-level model overrides - Pass appropriate models to classifier vs feedback functions This ensures the CLI simulation tool matches the production activity.py behavior. --- research/guarded_ai.py | 73 +++++++++++++++++++++++++++++++----------- 1 file changed, 55 insertions(+), 18 deletions(-) diff --git a/research/guarded_ai.py b/research/guarded_ai.py index 4a3054b..228ea1c 100644 --- a/research/guarded_ai.py +++ b/research/guarded_ai.py @@ -36,11 +36,35 @@ def initialize_model_map(): def get_openai_client_and_model(model_name=None): - """Get OpenAI client and model name""" - if not model_name: - model_name = "adamo1139/Hermes-3-Llama-3.1-8B-FP8-Dynamic" + """Get OpenAI client and model name - # Try to find client for specific model + Supports both direct model names and MODEL_X environment variable references. + If model_name is MODEL_1, MODEL_2, etc., looks up from environment. + """ + # Handle MODEL_X references + if model_name and model_name.startswith("MODEL_"): + # Extract the number from MODEL_X + try: + model_num = model_name.split("_")[1] + endpoint_key = f"MODEL_ENDPOINT_{model_num}" + api_key_key = f"MODEL_API_KEY_{model_num}" + + endpoint = os.getenv(endpoint_key) + api_key = os.getenv(api_key_key) + + if endpoint and api_key: + client = get_client_for_endpoint(endpoint, api_key) + # Use a simple default model name for the endpoint + actual_model = "model" # Most endpoints use "model" or ignore this + return client, actual_model + except Exception as e: + print(f"Warning: Failed to load {model_name}: {e}, falling back to default") + + # Default to MODEL_1 (Hermes) + if not model_name: + return get_openai_client_and_model("MODEL_1") + + # Try to find client for specific model name for stored_model, (client, base_url) in MODEL_CLIENT_MAP.items(): if model_name in stored_model or stored_model == model_name: return client, model_name @@ -68,8 +92,8 @@ def load_yaml_activity(file_path): return yaml.safe_load(file) -# Categorize the user's response using gpt-4o-mini -def categorize_response(question, response, buckets, tokens_for_ai): +# Categorize the user's response +def categorize_response(question, response, buckets, tokens_for_ai, model="MODEL_1"): bucket_list = ", ".join([str(bucket) for bucket in buckets]) messages = [ { @@ -83,7 +107,7 @@ def categorize_response(question, response, buckets, tokens_for_ai): ] try: - client, model_name = get_openai_client_and_model() + client, model_name = get_openai_client_and_model(model) completion = client.chat.completions.create( model=model_name, messages=messages, @@ -98,8 +122,8 @@ def categorize_response(question, response, buckets, tokens_for_ai): return f"Error: {e}" -# Generate AI feedback using gpt-4o-mini -def generate_ai_feedback(category, question, user_response, tokens_for_ai, metadata): +# Generate AI feedback +def generate_ai_feedback(category, question, user_response, tokens_for_ai, metadata, model="MODEL_1"): messages = [ { "role": "system", @@ -112,7 +136,7 @@ def generate_ai_feedback(category, question, user_response, tokens_for_ai, metad ] try: - client, model_name = get_openai_client_and_model() + client, model_name = get_openai_client_and_model(model) completion = client.chat.completions.create( model=model_name, messages=messages, max_tokens=250, temperature=0.7 ) @@ -131,6 +155,7 @@ def provide_feedback( user_language, tokens_for_ai, metadata, + model="MODEL_1", ): feedback = "" if "ai_feedback" in transition: @@ -143,7 +168,7 @@ def provide_feedback( feedback_metadata = {k: v for k, v in metadata.items() if k in filter_keys} ai_feedback = generate_ai_feedback( - category, question, user_response, tokens_for_ai, feedback_metadata + category, question, user_response, tokens_for_ai, feedback_metadata, model ) feedback += f"\n\nAI Feedback: {ai_feedback}" @@ -160,6 +185,7 @@ def provide_feedback_prompts( user_language, metadata, legacy_tokens_for_ai="", + model="MODEL_1", ): """Generate feedback from multiple prompts""" feedback_messages = [] @@ -200,7 +226,7 @@ def provide_feedback_prompts( filtered_user_response = "" # Remove user response if not in filter ai_feedback = generate_ai_feedback( - category, question, filtered_user_response, tokens_for_ai, prompt_metadata + category, question, filtered_user_response, tokens_for_ai, prompt_metadata, model ) # Only add feedback if it has content and isn't exactly the STFU token @@ -246,7 +272,7 @@ def get_next_section_and_step(activity_content, current_section_id, current_step return None, None -def translate_text(text, target_language): +def translate_text(text, target_language, model="MODEL_1"): # Guard clause for default language if target_language.lower() == "english": return text @@ -263,8 +289,9 @@ def translate_text(text, target_language): ] try: + client, model_name = get_openai_client_and_model(model) completion = client.chat.completions.create( - model="gpt-4o-mini", messages=messages, max_tokens=500, temperature=0.7 + model=model_name, messages=messages, max_tokens=500, temperature=0.7 ) translation = completion.choices[0].message.content.strip() return translation @@ -276,6 +303,10 @@ def simulate_activity(yaml_file_path): yaml_content = load_yaml_activity(yaml_file_path) max_attempts = yaml_content.get("default_max_attempts_per_step", 3) + # Get activity-level model defaults (default to MODEL_1 - Hermes) + default_classifier_model = yaml_content.get("classifier_model", "MODEL_1") + default_feedback_model = yaml_content.get("feedback_model", "MODEL_1") + current_section_id = yaml_content["sections"][0]["section_id"] current_step_id = yaml_content["sections"][0]["steps"][0]["step_id"] @@ -298,13 +329,17 @@ def simulate_activity(yaml_file_path): (s for s in section["steps"] if s["step_id"] == current_step_id), None ) + # Get step-level model overrides (if specified), otherwise use activity defaults + classifier_model = step.get("classifier_model", default_classifier_model) + feedback_model = step.get("feedback_model", default_feedback_model) + # Get the user's language preference from metadata user_language = metadata.get("language", "English") # Translate and print all content blocks once per step if "content_blocks" in step: content = "\n\n".join(step["content_blocks"]) - translated_content = translate_text(content, user_language) + translated_content = translate_text(content, user_language, feedback_model) print(translated_content) # Skip classification and feedback if there's no question @@ -315,7 +350,7 @@ def simulate_activity(yaml_file_path): continue question = step["question"] - translated_question = translate_text(question, user_language) + translated_question = translate_text(question, user_language, feedback_model) print(f"\nQuestion: {translated_question}") attempts = 0 @@ -338,7 +373,7 @@ def simulate_activity(yaml_file_path): print(f"DEBUG: Pre-script completed, updated metadata") category = categorize_response( - question, user_response, step["buckets"], step["tokens_for_ai"] + question, user_response, step["buckets"], step["tokens_for_ai"], classifier_model ) print(f"\nCategory: {category}") @@ -377,7 +412,7 @@ def simulate_activity(yaml_file_path): if "content_blocks" in transition: transition_content = "\n\n".join(transition["content_blocks"]) translated_transition_content = translate_text( - transition_content, user_language + transition_content, user_language, feedback_model ) print(translated_transition_content) @@ -490,6 +525,7 @@ def simulate_activity(yaml_file_path): step.get( "feedback_tokens_for_ai", "" ), # Pass legacy tokens to be combined + feedback_model, ) feedback_messages.extend(multi_feedback_messages) elif step.get("feedback_tokens_for_ai"): @@ -502,6 +538,7 @@ def simulate_activity(yaml_file_path): user_language, step.get("feedback_tokens_for_ai", ""), metadata, + feedback_model, ) if feedback and feedback.strip(): feedback_messages.append({"name": "Feedback", "content": feedback})