{
  "cells": [
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "# 🦙 DigiFair AI — Llama / LoRA Training with W&B\n",
        "\n",
        "این نوت‌بوک برای اضافه کردن فرایند آموزش/فاین‌تیون به پروژه Exhibition RAG ساخته شده است.\n",
        "\n",
        "✅ ویژگی‌ها:\n",
        "- ساخت دیتاست آموزشی از شرکت‌های نمایشگاه\n",
        "- آموزش سبک با LoRA/QLoRA\n",
        "- لاگ کردن loss و metric در Weights & Biases\n",
        "- ذخیره adapter خروجی\n",
        "- امکان push به Hugging Face Hub\n",
        "\n",
        "⚠️ نکته امنیتی:\n",
        "- `WANDB_API_KEY` و `HF_TOKEN` را فقط در Colab Secrets بگذارید.\n",
        "- هیچ توکنی را داخل سلول‌ها hard-code نکنید.\n"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## 0) انتخاب مدل\n",
        "\n",
        "اگر به مدل‌های رسمی Meta Llama دسترسی داری، می‌توانی `MODEL_ID` را روی یکی از مدل‌های Llama بگذاری.\n",
        "اگر دسترسی نداری، برای تست آموزشی از TinyLlama استفاده کن.\n"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {},
      "outputs": [],
      "source": [
        "# برای تست سبک و عمومی:\n",
        "MODEL_ID = \"TinyLlama/TinyLlama-1.1B-Chat-v1.0\"\n",
        "\n",
        "# اگر HF_TOKEN و دسترسی Llama داری، می‌توانی یکی از این‌ها را امتحان کنی:\n",
        "# MODEL_ID = \"meta-llama/Llama-3.2-1B-Instruct\"\n",
        "# MODEL_ID = \"meta-llama/Meta-Llama-3.1-8B-Instruct\"\n",
        "\n",
        "PROJECT_NAME = \"exhibition-rag-llama-training\"\n",
        "RUN_NAME = \"digifair-lora-run\"\n",
        "OUTPUT_DIR = \"/content/digifair-lora-adapter\"\n"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## 1) نصب کتابخانه‌ها"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {},
      "outputs": [],
      "source": [
        "!pip install -q -U transformers accelerate datasets peft trl bitsandbytes wandb huggingface_hub pandas requests scikit-learn\n"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## 2) اتصال امن به Hugging Face و W&B"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {},
      "outputs": [],
      "source": [
        "import os\n",
        "from getpass import getpass\n",
        "\n",
        "def get_secret(name, prompt):\n",
        "    # Colab Secrets first\n",
        "    try:\n",
        "        from google.colab import userdata\n",
        "        val = userdata.get(name)\n",
        "        if val:\n",
        "            return val\n",
        "    except Exception:\n",
        "        pass\n",
        "    # Environment\n",
        "    val = os.environ.get(name)\n",
        "    if val:\n",
        "        return val\n",
        "    # Manual input\n",
        "    return getpass(prompt)\n",
        "\n",
        "HF_TOKEN = get_secret(\"HF_TOKEN\", \"HF_TOKEN را وارد کنید یا Enter بزنید اگر لازم نیست: \")\n",
        "WANDB_API_KEY = get_secret(\"WANDB_API_KEY\", \"WANDB_API_KEY را وارد کنید یا Enter بزنید اگر نمی‌خواهید W&B فعال شود: \")\n",
        "\n",
        "if HF_TOKEN:\n",
        "    os.environ[\"HF_TOKEN\"] = HF_TOKEN\n",
        "if WANDB_API_KEY:\n",
        "    os.environ[\"WANDB_API_KEY\"] = WANDB_API_KEY\n",
        "    os.environ[\"WANDB_PROJECT\"] = PROJECT_NAME\n",
        "    os.environ[\"WANDB_NAME\"] = RUN_NAME\n",
        "\n",
        "print(\"HF token:\", bool(HF_TOKEN))\n",
        "print(\"W&B token:\", bool(WANDB_API_KEY))\n"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## 3) دانلود دیتای شرکت‌ها"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {},
      "outputs": [],
      "source": [
        "import requests, pandas as pd, json, re\n",
        "\n",
        "COMPANIES_JSON_URL = \"https://sosa123456-exhibition-connector-rag1-static.static.hf.space/data/companies.json\"\n",
        "resp = requests.get(COMPANIES_JSON_URL, timeout=60)\n",
        "resp.raise_for_status()\n",
        "companies = resp.json()\n",
        "print(\"records:\", len(companies))\n",
        "companies[0]\n"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## 4) ساخت دیتاست Instruction برای Llama"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {},
      "outputs": [],
      "source": [
        "def clean(x):\n",
        "    return str(x or \"\").strip()\n",
        "\n",
        "def company_answer(r):\n",
        "    return f\"\"\"نام شرکت: {clean(r.get('company'))}\n",
        "زمینه فعالیت: {clean(r.get('activity')) or 'نامشخص'}\n",
        "زمینه کاری: {clean(r.get('category')) or 'نامشخص'}\n",
        "وب‌سایت: {clean(r.get('website')) or 'ثبت نشده'}\n",
        "سالن/غرفه: {clean(r.get('hall')) or 'ثبت نشده'} / {clean(r.get('booth')) or 'ثبت نشده'}\"\"\"\n",
        "\n",
        "examples = []\n",
        "for r in companies:\n",
        "    c = clean(r.get('company'))\n",
        "    if not c or c == 'بدون نام':\n",
        "        continue\n",
        "    examples.append({\n",
        "        \"instruction\": f\"اطلاعات شرکت {c} را در نمایشگاه توضیح بده.\",\n",
        "        \"output\": company_answer(r)\n",
        "    })\n",
        "    if r.get('activity'):\n",
        "        examples.append({\n",
        "            \"instruction\": f\"کدام شرکت در زمینه {clean(r.get('activity'))[:80]} فعالیت می‌کند؟\",\n",
        "            \"output\": company_answer(r)\n",
        "        })\n",
        "    if r.get('hall'):\n",
        "        examples.append({\n",
        "            \"instruction\": f\"کدام شرکت در سالن {clean(r.get('hall'))} حضور دارد؟\",\n",
        "            \"output\": company_answer(r)\n",
        "        })\n",
        "\n",
        "# محدود کردن برای آموزش سریع دمو\n",
        "examples = examples[:2500]\n",
        "print(\"training examples:\", len(examples))\n",
        "pd.DataFrame(examples).head()\n"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {},
      "outputs": [],
      "source": [
        "from datasets import Dataset\n",
        "\n",
        "def format_prompt(row):\n",
        "    return f\"\"\"<|system|>\n",
        "شما دستیار فارسی نمایشگاه هستید. فقط بر اساس داده‌های نمایشگاه پاسخ بدهید.\n",
        "<|user|>\n",
        "{row['instruction']}\n",
        "<|assistant|>\n",
        "{row['output']}\"\"\"\n",
        "\n",
        "dataset = Dataset.from_list([{\"text\": format_prompt(x)} for x in examples])\n",
        "dataset = dataset.train_test_split(test_size=0.05, seed=42)\n",
        "dataset\n"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## 5) بارگذاری مدل و آماده‌سازی LoRA"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {},
      "outputs": [],
      "source": [
        "import torch\n",
        "from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig\n",
        "from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training\n",
        "\n",
        "use_cuda = torch.cuda.is_available()\n",
        "print(\"CUDA:\", use_cuda, torch.cuda.get_device_name(0) if use_cuda else \"CPU\")\n",
        "\n",
        "tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, token=HF_TOKEN or None, trust_remote_code=True)\n",
        "if tokenizer.pad_token is None:\n",
        "    tokenizer.pad_token = tokenizer.eos_token\n",
        "\n",
        "if use_cuda:\n",
        "    bnb_config = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16, bnb_4bit_quant_type=\"nf4\")\n",
        "    model = AutoModelForCausalLM.from_pretrained(MODEL_ID, token=HF_TOKEN or None, quantization_config=bnb_config, device_map=\"auto\", trust_remote_code=True)\n",
        "    model = prepare_model_for_kbit_training(model)\n",
        "else:\n",
        "    model = AutoModelForCausalLM.from_pretrained(MODEL_ID, token=HF_TOKEN or None, trust_remote_code=True)\n",
        "\n",
        "lora_config = LoraConfig(\n",
        "    r=16, lora_alpha=32, lora_dropout=0.05, bias=\"none\", task_type=\"CAUSAL_LM\",\n",
        "    target_modules=[\"q_proj\", \"v_proj\", \"k_proj\", \"o_proj\"] if \"llama\" in MODEL_ID.lower() else [\"q_proj\", \"v_proj\"]\n",
        ")\n",
        "model = get_peft_model(model, lora_config)\n",
        "model.print_trainable_parameters()\n"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## 6) آموزش با W&B Logging"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {},
      "outputs": [],
      "source": [
        "import wandb\n",
        "from trl import SFTTrainer, SFTConfig\n",
        "\n",
        "report_to = []\n",
        "if WANDB_API_KEY:\n",
        "    wandb.login(key=WANDB_API_KEY)\n",
        "    wandb.init(project=PROJECT_NAME, name=RUN_NAME, config={\"model\": MODEL_ID, \"examples\": len(examples)})\n",
        "    report_to = [\"wandb\"]\n",
        "\n",
        "training_args = SFTConfig(\n",
        "    output_dir=OUTPUT_DIR,\n",
        "    per_device_train_batch_size=1 if use_cuda else 1,\n",
        "    gradient_accumulation_steps=8,\n",
        "    learning_rate=2e-4,\n",
        "    num_train_epochs=1,\n",
        "    max_steps=80,  # برای دمو سریع؛ برای آموزش واقعی بیشتر کن\n",
        "    logging_steps=5,\n",
        "    save_steps=40,\n",
        "    eval_steps=40,\n",
        "    report_to=report_to,\n",
        "    fp16=use_cuda,\n",
        "    bf16=False,\n",
        "    max_seq_length=1024,\n",
        ")\n",
        "\n",
        "trainer = SFTTrainer(\n",
        "    model=model,\n",
        "    tokenizer=tokenizer,\n",
        "    train_dataset=dataset[\"train\"],\n",
        "    eval_dataset=dataset[\"test\"],\n",
        "    args=training_args,\n",
        "    dataset_text_field=\"text\",\n",
        ")\n",
        "\n",
        "trainer.train()\n",
        "trainer.save_model(OUTPUT_DIR)\n",
        "if WANDB_API_KEY:\n",
        "    wandb.finish()\n",
        "print(\"saved to\", OUTPUT_DIR)\n"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## 7) تست مدل آموزش‌دیده"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {},
      "outputs": [],
      "source": [
        "from transformers import pipeline\n",
        "\n",
        "pipe = pipeline(\"text-generation\", model=model, tokenizer=tokenizer, max_new_tokens=180)\n",
        "prompt = \"\"\"<|system|>\n",
        "شما دستیار فارسی نمایشگاه هستید.\n",
        "<|user|>\n",
        "وب‌سایت شرکت پریسماتک چیست؟\n",
        "<|assistant|>\n",
        "\"\"\"\n",
        "out = pipe(prompt)[0][\"generated_text\"]\n",
        "print(out)\n"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## 8) آپلود Adapter به Hugging Face Hub اختیاری"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {},
      "outputs": [],
      "source": [
        "# اختیاری: اگر می‌خواهی adapter را روی Hugging Face ذخیره کنی\n",
        "# از Colab Secrets یک HF_TOKEN با write permission لازم است.\n",
        "\n",
        "PUSH_TO_HUB = False\n",
        "REPO_ID = \"YOUR_USERNAME/digifair-rag-lora\"\n",
        "\n",
        "if PUSH_TO_HUB:\n",
        "    model.push_to_hub(REPO_ID, token=HF_TOKEN)\n",
        "    tokenizer.push_to_hub(REPO_ID, token=HF_TOKEN)\n",
        "    print(\"pushed to\", REPO_ID)\n",
        "else:\n",
        "    print(\"Push disabled. Adapter saved locally at\", OUTPUT_DIR)\n"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## 9) اتصال نتیجه آموزش به سایت\n",
        "\n",
        "در نسخه واقعی:\n",
        "1. Adapter یا مدل را روی Hugging Face ذخیره کنید.\n",
        "2. در Vercel Env این‌ها را بگذارید:\n",
        "   - `HF_TOKEN`\n",
        "   - `HF_MODEL` یا `TRAINED_ADAPTER_REPO`\n",
        "   - `WANDB_API_KEY`\n",
        "3. Endpoint پاسخ‌دهی می‌تواند به مدل آموزش‌دیده وصل شود.\n",
        "\n",
        "فعلاً سایت production از Vercel JSON RAG استفاده می‌کند و این نوت‌بوک برای آموزش/آزمایش مدل است.\n"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## 10) تست Gradio فارسی RTL\n",
        "\n",
        "این سلول برای تست نمایش صحیح فارسی در Gradio است.\n"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {},
      "outputs": [],
      "source": [
        "!pip install -q gradio\n",
        "import gradio as gr\n",
        "\n",
        "def format_fa_answer(text):\n",
        "    text = str(text or \"\").strip()\n",
        "    # جداسازی ساده خطوط و بهبود خوانایی فارسی\n",
        "    text = text.replace(\"|\", \"\\n- \")\n",
        "    return text\n",
        "\n",
        "def demo_answer(message, history):\n",
        "    # اگر مدل pipe ساخته شده باشد، از آن استفاده می‌کنیم؛ وگرنه پاسخ دمو\n",
        "    try:\n",
        "        prompt = f\"\"\"<|system|>\n",
        "شما دستیار فارسی نمایشگاه هستید. پاسخ کوتاه، دقیق و راست‌چین بنویسید.\n",
        "<|user|>\n",
        "{message}\n",
        "<|assistant|>\n",
        "\"\"\"\n",
        "        out = pipe(prompt)[0][\"generated_text\"]\n",
        "        ans = out.split(\"<|assistant|>\")[-1].strip()\n",
        "    except Exception:\n",
        "        ans = \"نمونه پاسخ فارسی:\\n- شرکت پریسماتک در حوزه ابزار دقیق فعالیت دارد.\\n- وب‌سایت: https://prismatech.ir/\\n- اگر سالن/غرفه در اکسل ثبت نشده باشد، باید در پنل ادمین تکمیل شود.\"\n",
        "    return format_fa_answer(ans)\n",
        "\n",
        "css = \"\"\"\n",
        ".gradio-container { direction: rtl; font-family: Tahoma, Arial, sans-serif; }\n",
        "textarea, input, .prose, .markdown, .message { direction: rtl !important; text-align: right !important; }\n",
        "footer { display: none !important; }\n",
        "\"\"\"\n",
        "\n",
        "demo = gr.ChatInterface(\n",
        "    fn=demo_answer,\n",
        "    title=\"دستیار فارسی نمایشگاه — تست Gradio RTL\",\n",
        "    description=\"برای تست نوشتار فارسی، راست‌چین و پاسخ مدل آموزش‌دیده/دمو\",\n",
        "    examples=[\"وب‌سایت شرکت پریسماتک چیست؟\", \"مواد شیمیایی تصفیه آب\", \"سالن 31B\"],\n",
        "    css=css,\n",
        "    chatbot=gr.Chatbot(rtl=True, height=420, show_copy_button=True),\n",
        "    textbox=gr.Textbox(placeholder=\"سؤال خود را فارسی بنویسید...\", rtl=True),\n",
        ")\n",
        "\n",
        "demo.launch(share=True, debug=True)\n"
      ]
    }
  ],
  "metadata": {
    "colab": {
      "provenance": [],
      "include_colab_link": true
    },
    "kernelspec": {
      "display_name": "Python 3",
      "name": "python3"
    },
    "language_info": {
      "name": "python"
    }
  },
  "nbformat": 4,
  "nbformat_minor": 5
}