Spaces:
Sleeping
Sleeping
| import streamlit as st | |
| import pandas as pd | |
| import faiss | |
| import numpy as np | |
| from sentence_transformers import SentenceTransformer | |
| from transformers import AutoModelForCausalLM, AutoTokenizer, pipeline, AutoModelForSeq2SeqLM | |
| from PIL import Image | |
| from io import BytesIO | |
| import requests | |
| import torch | |
| import os | |
| # --- 1. Load Recipes Dataset --- | |
| def load_recipes(): | |
| try: | |
| recipes_df = pd.read_csv("recipes.csv") | |
| recipes_df = recipes_df.rename(columns={"recipe_name": "title", "directions": "instructions"}) | |
| recipes_df = recipes_df[['title', 'ingredients', 'instructions', 'img_src']] | |
| recipes_df.fillna("", inplace=True) | |
| recipes_df["ingredients"] = recipes_df["ingredients"].str.lower() | |
| recipes_df["combined_text"] = recipes_df["title"] + " " + recipes_df["ingredients"] | |
| return recipes_df | |
| except Exception as e: | |
| st.error(f"β Error loading recipes: {e}") | |
| return pd.DataFrame() | |
| recipes_df = load_recipes() | |
| # --- 2. Load SentenceTransformer Model --- | |
| def load_embedding_model(): | |
| return SentenceTransformer("all-MiniLM-L6-v2") | |
| embedding_model = load_embedding_model() | |
| # --- 3. Initialize FAISS Index --- | |
| FAISS_INDEX_PATH = "faiss_index.idx" | |
| def build_faiss_index(): | |
| dim = 384 | |
| index = faiss.IndexFlatL2(dim) | |
| if os.path.exists(FAISS_INDEX_PATH): | |
| index = faiss.read_index(FAISS_INDEX_PATH) | |
| else: | |
| combined_texts = recipes_df["combined_text"].tolist() | |
| embeddings = embedding_model.encode(combined_texts, normalize_embeddings=True, convert_to_numpy=True).astype(np.float32) | |
| index.add(embeddings) | |
| faiss.write_index(index, FAISS_INDEX_PATH) | |
| return index | |
| faiss_index = build_faiss_index() | |
| # --- 4. Load Classification Model (for query classification) --- | |
| def load_classifier(): | |
| return pipeline("zero-shot-classification", model="facebook/bart-large-mnli") | |
| classifier = load_classifier() | |
| # --- Load T5 Model (for sequence-to-sequence tasks and recipe generation) --- | |
| def load_llm_model(): | |
| model_name = "google/flan-t5-small" # T5-based model for sequence-to-sequence tasks | |
| tokenizer = AutoTokenizer.from_pretrained(model_name) | |
| model = AutoModelForSeq2SeqLM.from_pretrained(model_name, torch_dtype=torch.float32, device_map="cpu") # Use AutoModelForSeq2SeqLM | |
| return model, tokenizer | |
| llm_model, tokenizer = load_llm_model() | |
| # --- 6. Query Classification --- | |
| def classify_query(query): | |
| labels = ["recipe", "general", "ingredients"] | |
| result = classifier(query, candidate_labels=labels) | |
| return result['labels'][0] | |
| # --- 7. Recipe Retrieval --- | |
| def retrieve_recipes(query, top_k=3): | |
| query_embedding = embedding_model.encode(query, convert_to_numpy=True).reshape(1, -1) | |
| distances, indices = faiss_index.search(query_embedding, top_k) | |
| return recipes_df.iloc[indices[0]] if len(indices[0]) > 0 else None | |
| # --- 8. Generate a Recipe (Recipe Generation with T5-based model) --- | |
| def generate_recipe(ingredients): | |
| system_prompt = """You are a creative AI chef. Create a recipe based on the following ingredients: {ingredients}. Be sure to include instructions and specify the cuisine type.""" | |
| full_prompt = system_prompt.format(ingredients=ingredients) | |
| inputs = tokenizer(full_prompt, return_tensors="pt", truncation=True).to("cpu") | |
| with torch.no_grad(): | |
| output = llm_model.generate(**inputs, max_length=300, num_beams=2, early_stopping=True) | |
| generated_recipe = tokenizer.decode(output[0], skip_special_tokens=True) | |
| return generated_recipe.strip() | |
| # --- 9. Display Image --- | |
| def display_image(image_url, recipe_name): | |
| try: | |
| response = requests.get(image_url, timeout=5) | |
| response.raise_for_status() | |
| image = Image.open(BytesIO(response.content)) | |
| st.image(image, caption=recipe_name, use_container_width=True) | |
| except requests.exceptions.RequestException: | |
| st.image("https://via.placeholder.com/300?text=No+Image", caption=recipe_name, use_container_width=True) | |
| # --- 10. Handle Recipe Queries (Retrieve or Generate) --- | |
| def handle_recipe_query(query): | |
| retrieved_recipes = retrieve_recipes(query, top_k=3) | |
| if retrieved_recipes is not None and not retrieved_recipes.empty: | |
| st.subheader("π΄ Found Recipes:") | |
| for _, recipe in retrieved_recipes.iterrows(): | |
| st.markdown(f"### {recipe['title']}") | |
| st.write(f"**Ingredients:** {recipe['ingredients']}") | |
| st.write(f"**Instructions:** {recipe['instructions']}") | |
| display_image(recipe.get('img_src', ''), recipe['title']) | |
| else: | |
| st.write("β No recipes found. Generating a new recipe based on your ingredients...") | |
| generated_recipe = generate_recipe(query) | |
| st.subheader("π³ Generated Recipe:") | |
| st.write(generated_recipe) | |
| # --- 11. Answer General Queries --- | |
| def answer_question(query): | |
| system_prompt = """You are a helpful AI assistant. Provide clear and concise responses to user questions.""" | |
| full_prompt = f"{system_prompt}\nUser Query: {query}" | |
| inputs = tokenizer(full_prompt, return_tensors="pt", truncation=True).to("cpu") | |
| with torch.no_grad(): | |
| output = llm_model.generate(**inputs, max_length=200, num_beams=2, early_stopping=True) | |
| response = tokenizer.decode(output[0], skip_special_tokens=True) | |
| return response.strip() | |
| # --- 12. Streamlit UI --- | |
| st.title("π½οΈ AI Recipe Assistant") | |
| user_query = st.text_input('Enter your Ingredient or recipe search query (eg: "what can i make with strawberries? or Give me the recipe of pizza"):', "", key="main_query_input") | |
| if st.button("Ask AI"): | |
| if user_query: | |
| with st.spinner("π Searching..."): | |
| query_type = classify_query(user_query) | |
| if query_type == "recipe" or query_type == "ingredients": | |
| handle_recipe_query(user_query) | |
| else: | |
| st.subheader("π€ AI Answer:") | |
| ai_answer = answer_question(user_query) | |
| st.write(ai_answer) | |
| else: | |
| st.warning("β Please enter a query before clicking 'Ask AI'.") | |