TaruniSwathi/Java-CSharp-CodeGen
0
1import re2 3import gradio as gr4import spaces5import torch6from transformers import AutoModelForCausalLM, AutoTokenizer7 8 9# ============================================================10# MODEL CONFIGURATION11# ============================================================12 13MODEL_ID = "shibsankardhara2/Qwen2.5-Coder-1.5B-Java-CSharp_V5"14 15 16# ============================================================17# LOAD TOKENIZER18# ============================================================19 20print("Loading tokenizer...")21 22tokenizer = AutoTokenizer.from_pretrained(23 MODEL_ID,24 trust_remote_code=True,25)26 27if tokenizer.pad_token_id is None:28 tokenizer.pad_token_id = tokenizer.eos_token_id29 30 31# ============================================================32# LOAD MODEL ON CPU33# ============================================================34 35print("Loading model on CPU...")36 37model = AutoModelForCausalLM.from_pretrained(38 MODEL_ID,39 torch_dtype=torch.float16,40 low_cpu_mem_usage=True,41 trust_remote_code=True,42)43 44model.eval()45 46print("Model loaded successfully.")47 48 49# ============================================================50# PROMPT FUNCTIONS51# ============================================================52 53def add_java_hint(instruction: str) -> str:54 """55 Explicitly mention Java when it is missing from the56 natural-language instruction.57 """58 instruction = instruction.strip()59 60 if "java" in instruction.lower():61 return instruction62 63 return f"{instruction} Write the solution in Java."64 65 66def build_prompt(task: str, user_input: str) -> str:67 """68 Construct the prompt expected by the fine-tuned model.69 """70 user_input = user_input.strip()71 72 if task == "Natural Language → Java":73 return (74 "### Instruction:\n\n"75 f"{add_java_hint(user_input)}\n\n"76 "### Response:\n\n"77 )78 79 return (80 "### Instruction:\n\n"81 "Translate the following Java code into equivalent C#. "82 "Write only the C# solution.\n\n"83 "### Java:\n\n"84 f"{user_input}\n\n"85 "### Response:\n\n"86 )87 88 89# ============================================================90# GENERAL OUTPUT CLEANING91# ============================================================92 93def clean_output(generated_text: str) -> str:94 """95 Remove Markdown fences, repeated prompt sections and96 model-specific special tokens.97 """98 if not generated_text:99 return ""100 101 cleaned_text = generated_text.strip()102 103 # Extract the contents of a Markdown code block, when present.104 fenced_match = re.search(105 r"```(?:java|csharp|cs|c#)?\s*(.*?)```",106 cleaned_text,107 flags=re.DOTALL | re.IGNORECASE,108 )109 110 if fenced_match:111 cleaned_text = fenced_match.group(1).strip()112 113 # Remove prompt sections that the model may generate again.114 stop_markers = [115 "### Instruction:",116 "### Instruction\n",117 "### Java:",118 "### Java\n",119 "### Response:",120 "### Response\n",121 "<|im_start|>",122 "<|im_end|>",123 "<|endoftext|>",124 ]125 126 for marker in stop_markers:127 marker_position = cleaned_text.find(marker)128 129 if marker_position != -1:130 cleaned_text = cleaned_text[:marker_position].strip()131 132 # Remove remaining opening or closing fences.133 cleaned_text = re.sub(134 r"^```(?:java|csharp|cs|c#)?\s*",135 "",136 cleaned_text,137 flags=re.IGNORECASE,138 )139 140 cleaned_text = re.sub(141 r"\s*```$",142 "",143 cleaned_text,144 )145 146 return cleaned_text.strip()147 148 149# ============================================================150# C# OUTPUT CLEANING151# ============================================================152 153def clean_csharp_output(text: str) -> str:154 """155 Remove invalid virtual or override modifiers from standalone156 C# method snippets.157 158 If the output contains a class, struct, interface, record or159 enum declaration, modifiers are preserved.160 """161 if not text:162 return ""163 164 # Explicitly initialise code before it is accessed.165 code = text.strip()166 167 has_type_declaration = re.search(168 r"\b(class|struct|interface|record|enum)\b",169 code,170 flags=re.IGNORECASE,171 )172 173 if has_type_declaration:174 return code175 176 # public virtual int Method() -> public int Method()177 # protected override void Method() -> protected void Method()178 code = re.sub(179 r"\b(public|private|protected|internal)\s+"180 r"(?:virtual|override)\s+",181 r"\1 ",182 code,183 flags=re.IGNORECASE,184 )185 186 # virtual int Method() -> int Method()187 # override void Method() -> void Method()188 code = re.sub(189 r"(^|\n)(\s*)(?:virtual|override)\s+",190 r"\1\2",191 code,192 flags=re.IGNORECASE,193 )194 195 return code.strip()196 197 198# ============================================================199# CODE GENERATION200# ============================================================201 202@spaces.GPU(duration=120)203def generate_code(task: str, user_input: str) -> str:204 """205 Generate Java from natural language or translate Java to C#.206 """207 if not user_input or not user_input.strip():208 return "Please enter a requirement or Java code."209 210 prompt = build_prompt(task, user_input)211 212 max_new_tokens = (213 300214 if task == "Natural Language → Java"215 else 400216 )217 218 try:219 if not torch.cuda.is_available():220 return "Generation failed: GPU is not available."221 222 # Move the model to the allocated ZeroGPU device.223 model.to("cuda")224 model.eval()225 226 tokenized_inputs = tokenizer(227 prompt,228 return_tensors="pt",229 truncation=True,230 max_length=2048,231 )232 233 tokenized_inputs = {234 key: value.to("cuda")235 for key, value in tokenized_inputs.items()236 }237 238 with torch.inference_mode():239 generated_ids = model.generate(240 **tokenized_inputs,241 max_new_tokens=max_new_tokens,242 do_sample=False,243 use_cache=True,244 pad_token_id=tokenizer.pad_token_id,245 eos_token_id=tokenizer.eos_token_id,246 )247 248 prompt_length = tokenized_inputs["input_ids"].shape[1]249 250 new_token_ids = generated_ids[0][prompt_length:]251 252 generated_text = tokenizer.decode(253 new_token_ids,254 skip_special_tokens=True,255 clean_up_tokenization_spaces=False,256 )257 258 # Always run the general output cleaner.259 output_code = clean_output(generated_text)260 261 # Run the C#-specific cleaner only for Java → C#.262 if task == "Java → C#":263 output_code = clean_csharp_output(output_code)264 markdown_language = "csharp"265 else:266 markdown_language = "java"267 268 if not output_code:269 return (270 "The model returned an empty response. "271 "Please try a more specific input."272 )273 274 return (275 f"```{markdown_language}\n"276 f"{output_code}\n"277 "```"278 )279 280 except Exception as error:281 return (282 "Generation failed: "283 f"{type(error).__name__}: {error}"284 )285 286 finally:287 # Return the model to CPU after the ZeroGPU request.288 try:289 model.to("cpu")290 except Exception as move_error:291 print(f"Could not move model to CPU: {move_error}")292 293 if torch.cuda.is_available():294 torch.cuda.empty_cache()295 296 297# ============================================================298# UPDATE TEXTBOX299# ============================================================300 301def update_input(task: str):302 """303 Update the input textbox for the selected task.304 """305 if task == "Natural Language → Java":306 return gr.update(307 label="Natural-language requirement",308 placeholder=(309 "Example: Write a Java method to check whether "310 "a number is prime."311 ),312 value="",313 )314 315 return gr.update(316 label="Java code",317 placeholder=(318 "Example:\n"319 "public static int factorial(int n) {\n"320 " int result = 1;\n"321 " for (int i = 2; i <= n; i++) {\n"322 " result *= i;\n"323 " }\n"324 " return result;\n"325 "}"326 ),327 value="",328 )329 330 331# ============================================================332# GRADIO INTERFACE333# ============================================================334 335with gr.Blocks(title="Java and C# CodeGen") as demo:336 gr.Markdown(337 """338# Java and C# CodeGen339 340Generate Java code from natural-language requirements or translate Java code341into equivalent C# using a fine-tuned Qwen2.5-Coder model.342"""343 )344 345 task = gr.Dropdown(346 choices=[347 "Natural Language → Java",348 "Java → C#",349 ],350 value="Natural Language → Java",351 label="Select task",352 )353 354 user_input = gr.Textbox(355 label="Natural-language requirement",356 placeholder=(357 "Example: Write a Java method to check whether "358 "a number is prime."359 ),360 lines=14,361 )362 363 generate_button = gr.Button(364 "Generate Code",365 variant="primary",366 )367 368 output = gr.Markdown(369 value="Generated code will appear here."370 )371 372 task.change(373 fn=update_input,374 inputs=task,375 outputs=user_input,376 )377 378 generate_button.click(379 fn=generate_code,380 inputs=[381 task,382 user_input,383 ],384 outputs=output,385 )386 387 gr.Examples(388 examples=[389 [390 "Natural Language → Java",391 (392 "Write a Java method to calculate factorial "393 "of a number using a loop."394 ),395 ],396 [397 "Natural Language → Java",398 "Write a Java method to reverse a string.",399 ],400 [401 "Natural Language → Java",402 (403 "Write a Java method to check whether "404 "a number is prime."405 ),406 ],407 [408 "Java → C#",409 """public static int factorial(int n) {410 int result = 1;411 412 for (int i = 2; i <= n; i++) {413 result *= i;414 }415 416 return result;417}""",418 ],419 [420 "Java → C#",421 """public boolean isEven(int n) {422 return n % 2 == 0;423}""",424 ],425 ],426 inputs=[427 task,428 user_input,429 ],430 )431 432 433# ============================================================434# START APPLICATION435# ============================================================436 437if __name__ == "__main__":438 demo.queue(439 default_concurrency_limit=1,440 max_size=10,441 ).launch()