psychologyphd/CodeLlama-7b-Text-to-SQL-MPS-FineTuned-V4
CodeLlama-7b Text-to-SQL-MPS-FineTuned-v4 (Fine-Tuned on Mac MPS)
Model Description
This is the 4th version of the fine-tuned CodeLlama-7b-hf specifically optimized for Text-to-SQL tasks. It was trained on a MacBook Pro M3 using MPS (Metal Performance Shaders) acceleration. This version demostrate that MPS fine tuning can achieve 71.0% accuracy based on Exact Match.
Origin & Adaptation
This project is adapted from the Microsoft "Generative AI for Beginners" Course (Chapter 18: Fine-tuning).
- Original Source: Generative AI for Beginners
- Old Version: psychologyphd/CodeLlama-7b-Text-to-SQL-mps-finetuned
- Modifications: tons of modifications to achieve this accuracy.
Evaluation Results
The model was evaluated on a held-out test set from the b-mc2/sql-create-context dataset.
Performance Notes:
- Contextual Understanding: The model shows strong performance in mapping natural language questions to complex SQL schemas provided in the context.
- Limitations: 71% accuracy indicates that while the model handles standard filter and aggregates well,
- eyeballing the results, the model struggles with table that needs to be created(index 10,39) possibly due to lack of such complex training examples,
- quotes (actually those cases should be considered correct because schema is not given),
- and applying correct function in selection (such as count).
Example Output:
- Correct Output:
🔹 Index: 0 🎯 Truth: SELECT home FROM tablename11 WHERE date = "16 april 2008" 🤖 Gen : select home from tablename11 where date = '16 april 2008' ---------------------------------------- 🔹 Index: 1 🎯 Truth: SELECT MAX(game) FROM tablename34 WHERE team = "celtics" AND highassists = "hedo türkoğlu (4)" 🤖 Gen : select max(game) from tablename34 where team = 'celtics' and highassists = 'hedo türkoğlu (4)' ---------------------------------------- 🔹 Index: 2 🎯 Truth: SELECT country FROM tablename17 WHERE score = 72 - 66 - 72 = 210 🤖 Gen : select country from tablename17 where score = 72 - 66 - 72 = 210 ---------------------------------------- 🔹 Index: 3 🎯 Truth: SELECT AVG(gold) FROM tablename66 WHERE sport = "athletics" AND silver > 42 🤖 Gen : select avg(gold) from tablename66 where sport = 'athletics' and silver > 42 ---------------------------------------- 🔹 Index: 4 🎯 Truth: SELECT club FROM tablename36 WHERE headcoach = "casemiro mior" 🤖 Gen : select club from tablename36 where headcoach = 'casemiro mior' ---------------------------------------- 🔹 Index: 5 🎯 Truth: SELECT COUNT(highpoints) FROM table231867386 WHERE record = "5-17" 🤖 Gen : select count(highpoints) from table231867386 where record = '5-17' ---------------------------------------- 🔹 Index: 6 🎯 Truth: SELECT date FROM tablename10 WHERE awayteam = "st kilda" 🤖 Gen : select date from tablename10 where awayteam = 'st kilda' ---------------------------------------- 🔹 Index: 8 🎯 Truth: SELECT sailnumber FROM table255952091 WHERE skipper = "Matt Allen" 🤖 Gen : select sailnumber from table255952091 where skipper = 'matt allen' ---------------------------------------- 🔹 Index: 9 🎯 Truth: SELECT venue FROM tablename74 WHERE hometeam = "melbourne" 🤖 Gen : select venue from tablename74 where hometeam = 'melbourne' ---------------------------------------- 🔹 Index: 11 🎯 Truth: SELECT time FROM tablename1 WHERE event = "k-1 the challenge 1999" 🤖 Gen : select time from tablename1 where event = 'k-1 the challenge 1999' ---------------------------------------- 🔹 Index: 12 🎯 Truth: SELECT COUNT(gold) FROM tablename34 WHERE silver = 2 AND total < 7 🤖 Gen : select count(gold) from tablename34 where silver = 2 and total < 7 ---------------------------------------- 🔹 Index: 13 🎯 Truth: SELECT player FROM tablename92 WHERE team = "chicago bulls" 🤖 Gen : select player from tablename92 where team = 'chicago bulls' ---------------------------------------- 🔹 Index: 14 🎯 Truth: SELECT player FROM table267906112 WHERE collegejuniorclubteam = "Litvinov (Czechoslovakia)" 🤖 Gen : select player from table267906112 where collegejuniorclubteam = 'litvinov (czechoslovakia)' ---------------------------------------- 🔹 Index: 15 🎯 Truth: SELECT DISTINCT name FROM instructor ORDER BY name 🤖 Gen : select distinct name from instructor order by name ---------------------------------------- 🔹 Index: 17 🎯 Truth: SELECT school FROM table116776912 WHERE college = "South Carolina" 🤖 Gen : select school from table116776912 where college = 'south carolina' ---------------------------------------- 🔹 Index: 19 🎯 Truth: SELECT venue FROM tablename47 WHERE score = "0–0" 🤖 Gen : select venue from tablename47 where score = '0–0' ---------------------------------------- 🔹 Index: 20 🎯 Truth: SELECT name FROM tablename88 WHERE nationality = "france" AND lane < 3 🤖 Gen : select name from tablename88 where nationality = 'france' and lane < 3 ---------------------------------------- 🔹 Index: 21 🎯 Truth: SELECT nation FROM tablename54 WHERE total < 19 AND bronze < 1 🤖 Gen : select nation from tablename54 where total < 19 and bronze < 1 ---------------------------------------- 🔹 Index: 22 🎯 Truth: SELECT record FROM tablename72 WHERE date = "october 27" 🤖 Gen : select record from tablename72 where date = 'october 27' ---------------------------------------- 🔹 Index: 23 🎯 Truth: SELECT score FROM tablename33 WHERE date = "october 17, 2007" 🤖 Gen : select score from tablename33 where date = 'october 17, 2007' ---------------------------------------- 🔹 Index: 24 🎯 Truth: SELECT nationality FROM tablename60 WHERE position = "forward" AND yearsforgrizzlies = "2011" 🤖 Gen : select nationality from tablename60 where position = 'forward' and yearsforgrizzlies = '2011' ---------------------------------------- 🔹 Index: 25 🎯 Truth: SELECT surface FROM tablename45 WHERE partner = "galina voskoboeva" 🤖 Gen : select surface from tablename45 where partner = 'galina voskoboeva' ---------------------------------------- 🔹 Index: 27 🎯 Truth: SELECT rank FROM tablename96 WHERE bronze < 7 AND nation = "norway" 🤖 Gen : select rank from tablename96 where bronze < 7 and nation = 'norway' ---------------------------------------- 🔹 Index: 28 🎯 Truth: SELECT sanskrt FROM tablename38 WHERE japanese = "jayana" 🤖 Gen : select sanskrt from tablename38 where japanese = 'jayana' ---------------------------------------- 🔹 Index: 32 🎯 Truth: SELECT format FROM tablename74 WHERE type = "primary" AND callletters = "kbjs" 🤖 Gen : select format from tablename74 where type = 'primary' and callletters = 'kbjs' ---------------------------------------- 🔹 Index: 34 🎯 Truth: SELECT MIN(byes) FROM tablename3 WHERE against = 1946 AND wins > 2 🤖 Gen : select min(byes) from tablename3 where against = 1946 and wins > 2 ---------------------------------------- 🔹 Index: 36 🎯 Truth: SELECT date FROM tablename61 WHERE attendance = "79,431" 🤖 Gen : select date from tablename61 where attendance = '79,431' ---------------------------------------- 🔹 Index: 38 🎯 Truth: SELECT MIN(manhuntinternational) FROM table300184601 🤖 Gen : select min(manhuntinternational) from table300184601 ---------------------------------------- 🔹 Index: 40 🎯 Truth: SELECT COUNT() FROM device 🤖 Gen : select count() from device ---------------------------------------- 🔹 Index: 42 🎯 Truth: SELECT COUNT(played) FROM tablename42 WHERE position < 4 AND team = "witton albion" 🤖 Gen : select count(played) from tablename42 where position < 4 and team = 'witton albion' ---------------------------------------- 🔹 Index: 43 🎯 Truth: SELECT score FROM tablename98 WHERE loss = "embree (1-2)" 🤖 Gen : select score from tablename98 where loss = 'embree (1-2)' ---------------------------------------- 🔹 Index: 44 🎯 Truth: SELECT dateofappointment FROM tablename91 WHERE mannerofdeparture = "sacked" 🤖 Gen : select dateofappointment from tablename91 where mannerofdeparture = 'sacked' ---------------------------------------- 🔹 Index: 45 🎯 Truth: SELECT catalog FROM tablename33 WHERE label = "grilled cheese" 🤖 Gen : select catalog from tablename33 where label = 'grilled cheese' ---------------------------------------- 🔹 Index: 46 🎯 Truth: SELECT medal FROM tablename48 WHERE games = "2008 beijing" AND sport = "taekwondo" 🤖 Gen : select medal from tablename48 where games = '2008 beijing' and sport = 'taekwondo' ---------------------------------------- 🔹 Index: 47 🎯 Truth: SELECT engine FROM tablename49 WHERE team = "arciero racing" 🤖 Gen : select engine from tablename49 where team = 'arciero racing' ---------------------------------------- 🔹 Index: 49 🎯 Truth: SELECT municipality FROM tablename53 WHERE station = "benson street" 🤖 Gen : select municipality from tablename53 where station = 'benson street' ----------------------------------------
- Wrong Output:
🔹 Index: 7 🎯 Truth: SELECT AVG(drawn) FROM tablename45 WHERE points1 = "33" AND goalsagainst = 45 AND played < 28 🤖 Gen : select avg(drawn) from tablename45 where points1 = 33 and goalsagainst = 45 and played < 28 ---------------------------------------- 🔹 Index: 10 🎯 Truth: SELECT major, age FROM student WHERE NOT stuid IN (SELECT T1.stuid FROM student AS T1 JOIN haspet AS T2 ON T1.stuid = T2.stuid JOIN pets AS T3 ON T3.petid = T2.petid WHERE T3.pettype = 'cat') 🤖 Gen : select t1.major, t1.age from student as t1 join haspet as t2 on t1.stuid = t2.stuid join pets as t3 on t2.petid = t3.petid where t3.pettype <> 'cat' ---------------------------------------- 🔹 Index: 16 🎯 Truth: SELECT COUNT(goals) FROM tablename49 WHERE goalratio < 0.8 AND games = 56 🤖 Gen : select goals from tablename49 where goalratio < 0.8 and games = 56 ---------------------------------------- 🔹 Index: 18 🎯 Truth: SELECT SUM(lost) FROM tablename86 WHERE points = 11 AND drawn > 1 🤖 Gen : select sum(lost) from tablename86 where points > 11 and drawn > 1 ---------------------------------------- 🔹 Index: 26 🎯 Truth: SELECT reward FROM table18932762 WHERE airdate = "October 6, 2005" 🤖 Gen : select count(reward) from table18932762 where airdate = 'october 6, 2005' ---------------------------------------- 🔹 Index: 29 🎯 Truth: SELECT name FROM tablename19 WHERE dcsfnumber = 3373 🤖 Gen : select name from tablename19 where dcsfnumber = '3373' ---------------------------------------- 🔹 Index: 30 🎯 Truth: SELECT COUNT(average) FROM table120946091 WHERE eveninggown = "8.988" 🤖 Gen : select count(average) from table120946091 where eveninggown = 8.988 ---------------------------------------- 🔹 Index: 31 🎯 Truth: SELECT writers FROM table178100993 WHERE episode_number = "14" 🤖 Gen : select writers from table178100993 where episodenumber = 14 ---------------------------------------- 🔹 Index: 33 🎯 Truth: SELECT location FROM tablename65 WHERE round = 3 AND method = "decision" 🤖 Gen : select location from tablename65 where round = '3' and method = 'won or loss by a decision' ---------------------------------------- 🔹 Index: 35 🎯 Truth: SELECT writtenby FROM table278233591 WHERE season_number = "23" 🤖 Gen : select writtenby from table278233591 where season_number = 23 ---------------------------------------- 🔹 Index: 37 🎯 Truth: SELECT attendance FROM tablename53 WHERE gamesite = "bye" 🤖 Gen : select count(attendance) from tablename53 where gamesite = 'bye' ---------------------------------------- 🔹 Index: 39 🎯 Truth: SELECT T1.fname, T1.lname FROM student AS T1 JOIN livesin AS T2 ON T1.stuid = T2.stuid WHERE T2.dormid IN (SELECT T3.dormid FROM hasamenity AS T3 JOIN dormamenity AS T4 ON T3.amenid = T4.amenid WHERE T4.amenityname = 'TV Lounge') 🤖 Gen : select t1.fname, t1.lname from student as t1 join livesin as t2 on t1.stuid = t2.stuid join hasamenity as t3 on t2.dormid = t3.dormid join dormamenity as t4 on t3.amenid = t4.amenid where t4.amenityname = 'tv lounge' group by t1.st ---------------------------------------- 🔹 Index: 41 🎯 Truth: SELECT attendance FROM tablename65 WHERE tieno = "4" 🤖 Gen : select count(attendance) from tablename65 where tieno = '4' ---------------------------------------- 🔹 Index: 48 🎯 Truth: SELECT hometeam AS score FROM tablename49 WHERE awayteam = "hawthorn" 🤖 Gen : select hometeam from tablename49 where away_team = 'hawthorn' ----------------------------------------
How to Use
See howtouse_v4.ipynb.
- mps pipeline actually works in fine tuning. For this light weighted how to use, I still use model.generate. .
Training Details
- Hardware: Mac M3 (MPS)
- Base Model: codellama/CodeLlama-7b-hf
- Dataset: b-mc2/sql-create-context
- Technique: LoRA (PEFT)
