{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# Purple Squirrel R1 - Demo Notebook\n", "\n", "This notebook demonstrates how to use the Purple Squirrel R1 model, a fine-tuned version of DeepSeek-R1-Distill-Llama-8B specialized for Purple Squirrel platform questions.\n", "\n", "[![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/purplesquirrelnetworks/purple-squirrel-r1/blob/main/purple_squirrel_r1_demo.ipynb)" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Option 1: Use the REST API (Recommended)\n", "\n", "The fastest way to use the model - no GPU required!" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import requests\n", "\n", "API_URL = \"https://purplesquirrelnetworks--purple-squirrel-r1-api-purplesqu-1e3c01.modal.run\"\n", "\n", "def ask_purple_squirrel(prompt, max_tokens=256, temperature=0.3):\n", " response = requests.post(\n", " API_URL,\n", " json={\"prompt\": prompt, \"max_tokens\": max_tokens, \"temperature\": temperature}\n", " )\n", " return response.json()[\"response\"]\n", "\n", "# Test it out\n", "print(ask_purple_squirrel(\"What is Purple Squirrel?\"))" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# More examples\n", "questions = [\n", " \"Can I mint NFTs with Purple Squirrel?\",\n", " \"What AI capabilities does Purple Squirrel have?\",\n", " \"How does Purple Squirrel handle video transcription?\",\n", "]\n", "\n", "for q in questions:\n", " print(f\"Q: {q}\")\n", " print(f\"A: {ask_purple_squirrel(q)}\")\n", " print(\"-\" * 50)" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Option 2: Load Model Locally (Requires GPU)\n", "\n", "For local inference with full control. Requires ~16GB VRAM or use 4-bit quantization." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# Install dependencies\n", "!pip install -q torch transformers accelerate bitsandbytes" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import torch\n", "from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig\n", "\n", "MODEL_ID = \"purplesquirrelnetworks/purple-squirrel-r1\"\n", "\n", "# 4-bit quantization for lower memory usage\n", "quantization_config = BitsAndBytesConfig(\n", " load_in_4bit=True,\n", " bnb_4bit_compute_dtype=torch.bfloat16,\n", " bnb_4bit_use_double_quant=True,\n", " bnb_4bit_quant_type=\"nf4\",\n", ")\n", "\n", "print(\"Loading tokenizer...\")\n", "tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)\n", "\n", "print(\"Loading model with 4-bit quantization...\")\n", "model = AutoModelForCausalLM.from_pretrained(\n", " MODEL_ID,\n", " quantization_config=quantization_config,\n", " device_map=\"auto\",\n", ")\n", "print(\"Model loaded!\")" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "def generate_response(prompt, max_tokens=256, temperature=0.3):\n", " messages = [{\"role\": \"user\", \"content\": prompt}]\n", " formatted = tokenizer.apply_chat_template(\n", " messages, tokenize=False, add_generation_prompt=True\n", " )\n", " inputs = tokenizer(formatted, return_tensors=\"pt\").to(model.device)\n", " \n", " outputs = model.generate(\n", " **inputs,\n", " max_new_tokens=max_tokens,\n", " temperature=temperature,\n", " do_sample=True,\n", " pad_token_id=tokenizer.eos_token_id,\n", " repetition_penalty=1.1,\n", " )\n", " \n", " response = tokenizer.decode(outputs[0], skip_special_tokens=True)\n", " # Extract assistant response\n", " if \"<|Assistant|>\" in response:\n", " response = response.split(\"<|Assistant|>\")[-1]\n", " if \"\" in response:\n", " response = response.split(\"\")[-1].strip()\n", " return response\n", "\n", "# Test\n", "print(generate_response(\"What is Purple Squirrel?\"))" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Model Info\n", "\n", "| Attribute | Value |\n", "|-----------|-------|\n", "| Base Model | DeepSeek-R1-Distill-Llama-8B |\n", "| Parameters | 8B |\n", "| Training Method | LoRA |\n", "| Training Examples | 74 |\n", "| Token Accuracy | 91% |\n", "\n", "## Links\n", "\n", "- [Model on HuggingFace](https://huggingface.co/purplesquirrelnetworks/purple-squirrel-r1)\n", "- [Purple Squirrel Networks](https://purplesquirrel.io)" ] } ], "metadata": { "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "name": "python", "version": "3.11.0" } }, "nbformat": 4, "nbformat_minor": 4 }