Views
No views yet
1%%capture
2# Installs Unsloth, Xformers (Flash Attention) and all other packages!
3!pip install "unsloth[colab-new] @ git+https://github.com/unslothai/unsloth.git"
4
5# We have to check which Torch version for Xformers (2.3 -> 0.0.27)
6from torch import __version__; from packaging.version import Version as V
7xformers = "xformers==0.0.27" if V(__version__) < V("2.4.0") else "xformers"
8!pip install --no-deps {xformers} trl peft accelerate bitsandbytes triton
9alpaca_prompt = """Below is an instruction that describes a task, paired with an input that provides further context. Write a response that appropriately completes the request.
10
11### Instruction:
12{}
13
14### Input:
15{}
16
17### Response:
18{}"""
19
20from unsloth import FastLanguageModel
21model, tokenizer = FastLanguageModel.from_pretrained(
22 model_name = "isaiahbjork/llama-3.1-8b-logic", # YOUR MODEL YOU USED FOR TRAINING
23 max_seq_length = max_seq_length,
24 dtype = dtype,
25 load_in_4bit = load_in_4bit,
26)
27FastLanguageModel.for_inference(model) # Enable native 2x faster inference
28
29
30inputs = tokenizer(
31[
32 alpaca_prompt.format(
33 "You are an expert at logic puzzles, reasoning, and planning", # instruction
34 "How many rs in strawberry?", # input
35 "", # output - leave this blank for generation!
36 )
37], return_tensors = "pt").to("cuda")
38
39from transformers import TextStreamer
40text_streamer = TextStreamer(tokenizer)
41_ = model.generate(**inputs, streamer = text_streamer, max_new_tokens = 256)1import re
2import random
3from transformers import TextStreamer
4
5# Function to parse the model output and extract the predicted count
6def extract_count(output):
7 # Make the regex pattern more flexible
8 match = re.search(r'(?:The letter "[a-z]"|\w+\'s) (?:appears?|occurs?|present?|is found|exists?) (\d+)', output, re.IGNORECASE)
9 if match:
10 return int(match.group(1))
11 return None
12
13# Function to generate test data
14def generate_test_data(num_words=150):
15 words = ["Airplane", "Airport", "Angelfish", "Antfarm", "Ballpark", "Beachball", "Bikerack", "Billboard", "Blackhole", "Blueberry", "Boardwalk", "Bodyguard", "Bookstore", "Bow Tie", "Brainstorm", "Busboy", "Cabdriver", "Candlestick", "Car wash", "Cartwheel", "Catfish", "Caveman", "Chocolate chip", "Crossbow", "Daydream", "Deadend", "Doghouse", "Dragonfly", "Dress shoes", "Dropdown", "Earlobe", "Earthquake", "Eyeballs", "Father-in-law", "Fingernail", "Firecracker", "Firefighter", "Firefly", "Firework", "Fishbowl", "Fisherman", "Fishhook", "Football", "Forget", "Forgive", "French fries", "Goodnight", "Grandchild", "Groundhog", "Hairband", "Hamburger", "Handcuff", "Handout", "Handshake", "Headband", "Herself", "High heels", "Honeydew", "Hopscotch", "Horseman", "Horseplay", "Hotdog", "Ice cream", "Itself", "Kickball", "Kickboxing", "Laptop", "Lifetime", "Lighthouse", "Mailman", "Midnight", "Milkshake", "Moonrocks", "Moonwalk", "Mother-in-law", "Movie theater", "Newborn", "Newsletter", "Newspaper", "Nightlight", "Nobody", "Northpole", "Nosebleed", "Outer space", "Over-the-counter", "Overestimate", "Paycheck", "Policeman", "Ponytail", "Post card", "Racquetball", "Railroad", "Rainbow", "Raincoat", "Raindrop", "Rattlesnake", "Rockband", "Rocketship", "Rowboat", "Sailboat", "Schoolbooks", "Schoolwork", "Shoelace", "Showoff", "Skateboard", "Snowball", "Snowflake", "Softball", "Solar system", "Soundproof", "Spaceship", "Spearmint", "Starfish", "Starlight", "Stingray", "Strawberry", "Subway", "Sunglasses", "Sunroof", "Supercharge", "Superman", "Superstar", "Tablespoon", "Tailbone", "Tailgate", "Take down", "Takeout", "Taxpayer", "Teacup", "Teammate", "Teaspoon", "Tennis shoes", "Throwback", "Timekeeper", "Timeline", "Timeshare", "Tugboat", "Tupperware", "Underestimate", "Uplift", "Upperclassman", "Uptown", "Video game", "Wallflower", "Waterboy", "Watermelon", "Wheelchair", "Without", "Workboots", "Worksheet"]
16
17 letters = "aeioulprts"
18 test_data = []
19 for word in words[:num_words]:
20 letter = random.choice(letters)
21 actual_count = word.lower().count(letter) # Use lower() to count case-insensitively
22 test_data.append((word, letter, actual_count))
23 return test_data
24
25# Alpaca prompt template
26alpaca_prompt = """Below is an instruction that describes a task, paired with an input that provides further context. Write a response that appropriately completes the request.
27
28### Instruction:
29{0}
30
31### Input:
32{1}
33
34### Response:
35"""
36
37# Generate test data
38test_data = generate_test_data()
39
40
41# Run evaluation
42correct_predictions = 0
43total_predictions = 0
44
45for word, letter, actual_count in test_data:
46 input_text = f"How many {letter}'s in {word}?"
47 prompt = alpaca_prompt.format(
48 "You are an expert at logic puzzles, reasoning, and planning",
49 input_text,
50 ""
51 )
52
53 inputs = tokenizer([prompt], return_tensors="pt").to("cuda")
54 text_streamer = TextStreamer(tokenizer)
55 output = model.generate(**inputs, streamer=text_streamer, max_new_tokens=256)
56
57 decoded_output = tokenizer.decode(output[0], skip_special_tokens=True)
58 print(f"Raw model output: {decoded_output}") # Print raw output for debugging
59 predicted_count = extract_count(decoded_output)
60
61 total_predictions += 1
62
63 if predicted_count is not None:
64 if predicted_count == actual_count:
65 correct_predictions += 1
66 else:
67 # If predicted_count is None and actual_count is 0, consider it correct
68 if actual_count == 0:
69 correct_predictions += 1
70 print(f"Warning: Could not extract a count from the model's response for '{word}'.")
71
72 print(f"Word: {word}, Letter: {letter}")
73 print(f"Actual count: {actual_count}, Predicted count: {predicted_count}")
74 print("Correct" if (predicted_count == actual_count or (predicted_count is None and actual_count == 0)) else "Incorrect")
75
76 # Calculate and print accuracy after each word
77 current_accuracy = correct_predictions / total_predictions
78 print(f"Current Accuracy: {current_accuracy:.2%}")
79 print(f"Correct predictions: {correct_predictions}")
80 print(f"Total predictions: {total_predictions}")
81 print("---")
82
83# Calculate accuracy
84accuracy = correct_predictions / total_predictions if total_predictions > 0 else 0
85print(f"\nAccuracy: {accuracy:.2%}")
86print(f"Correct predictions: {correct_predictions}")
87print(f"Total predictions: {total_predictions}")