nishu08/sql-codebert-classifier
013
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 