CoolFace
Modelpublic

nishu08/sql-codebert-classifier

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes13downloads
README.md66 linesDownload Raw Back to root
1---2language: en3license: mit4tags:5  - codebert6  - sql7  - education8  - text-classification9  - cross-encoder10base_model: microsoft/codebert-base11pipeline_tag: text-classification12---13 14# SQL CodeBERT Cross-Encoder15 16Multi-label SQL error classifier using **microsoft/codebert-base** as a cross-encoder.17 18## Input Format19 20All fields are concatenated into one sequence:21 22```23QUESTION:24{question}25 26SCHEMA:27{schema}28 29STUDENT_SQL:30{student_sql}31 32CORRECT_SQL:33{correct_sql}34```35 36## Labels37 38`JOIN_ERROR`, `AGGREGATION_ERROR`, `FILTER_ERROR`, `WINDOW_FUNCTION_ERROR`,39`SUBQUERY_ERROR`, `NULL_HANDLING_ERROR`, `PERFORMANCE_ERROR`, `LOGICAL_ERROR`, `SYNTAX_ERROR`40 41## Training42 43```bash44python -m src.hf_train_codebert \45  --data data/sql_errors_1m.parquet \46  --output-dir models/codebert-cross-encoder \47  --epochs 3 \48  --push-to-hub \49  --hub-model-id YOUR_USERNAME/sql-codebert-cross-encoder50```51 52## Inference53 54```python55from src.hf_predict_codebert import CodeBERTSQLErrorClassifier56 57clf = CodeBERTSQLErrorClassifier("YOUR_USERNAME/sql-codebert-cross-encoder")58result = clf.predict(59    question="What is the average score per department?",60    schema="students(id, score, department_id)",61    student_sql="SELECT department_id, SUM(score) FROM students GROUP BY department_id",62    correct_sql="SELECT department_id, AVG(score) FROM students GROUP BY department_id",63)64print(result["error_labels"])65```66