assix-research commited on
Commit
b7b1fb5
·
verified ·
1 Parent(s): a39ee6c

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +92 -0
app.py ADDED
@@ -0,0 +1,92 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import gradio as gr
2
+ import torch
3
+ import torch.nn as nn
4
+ from transformers import AutoTokenizer
5
+ from huggingface_hub import hf_hub_download
6
+
7
+ # 1. Model Architecture
8
+ class SourceCodeAuthorCheck(nn.Module):
9
+ def __init__(self, vocab_size=50257, d_model=128, nhead=8, num_layers=4, dim_feedforward=512):
10
+ super().__init__()
11
+ self.embedding = nn.Embedding(vocab_size, d_model)
12
+ self.pos_encoder = nn.Parameter(torch.zeros(1, 1024, d_model))
13
+
14
+ encoder_layers = nn.TransformerEncoderLayer(
15
+ d_model=d_model, nhead=nhead, dim_feedforward=dim_feedforward, batch_first=True
16
+ )
17
+ self.transformer = nn.TransformerEncoder(encoder_layers, num_layers=num_layers)
18
+ self.fc = nn.Linear(d_model, 1)
19
+
20
+ def forward(self, input_ids, attention_mask):
21
+ seq_len = input_ids.size(1)
22
+ x = self.embedding(input_ids) + self.pos_encoder[:, :seq_len, :]
23
+
24
+ src_key_padding_mask = ~attention_mask.bool()
25
+ x = self.transformer(x, src_key_padding_mask=src_key_padding_mask)
26
+
27
+ mask_expanded = attention_mask.unsqueeze(-1).float()
28
+ sum_embeddings = torch.sum(x * mask_expanded, 1)
29
+ sum_mask = torch.clamp(mask_expanded.sum(1), min=1e-9)
30
+ pooled = sum_embeddings / sum_mask
31
+
32
+ return self.fc(pooled)
33
+
34
+ # 2. Device and Loading Initialization
35
+ # Note: HF Free Spaces use CPU. We check for CUDA in case you upgrade the Space hardware.
36
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
37
+ tokenizer = AutoTokenizer.from_pretrained("gpt2")
38
+ tokenizer.pad_token = tokenizer.eos_token
39
+
40
+ model = SourceCodeAuthorCheck().to(device)
41
+
42
+ # Download weights securely from your model repository
43
+ model_path = hf_hub_download(repo_id="assix-research/SourceCodeAuthorCheck-SLM-10M", filename="source_code_classifier.pth")
44
+ model.load_state_dict(torch.load(model_path, map_location=device, weights_only=True))
45
+ model.eval()
46
+
47
+ # 3. Inference Logic
48
+ def predict_author(code_snippet):
49
+ if not code_snippet or not code_snippet.strip():
50
+ return "Please paste valid code.", "0.0%"
51
+
52
+ inputs = tokenizer(
53
+ code_snippet,
54
+ return_tensors="pt",
55
+ truncation=True,
56
+ padding="max_length",
57
+ max_length=1024
58
+ ).to(device)
59
+
60
+ with torch.no_grad():
61
+ # Handle mixed precision safely based on available hardware
62
+ if torch.cuda.is_available():
63
+ with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
64
+ logits = model(inputs['input_ids'], inputs['attention_mask'])
65
+ else:
66
+ logits = model(inputs['input_ids'], inputs['attention_mask'])
67
+
68
+ prob = torch.sigmoid(logits).item()
69
+
70
+ score = round(prob * 100, 2)
71
+ verdict = "🤖 AI Generated" if prob > 0.5 else "👨‍💻 Human Written"
72
+
73
+ return verdict, f"{score}%"
74
+
75
+ # 4. Gradio Interface Construction
76
+ demo = gr.Interface(
77
+ fn=predict_author,
78
+ inputs=gr.Code(language="python", label="Paste Python Source Code"),
79
+ outputs=[
80
+ gr.Textbox(label="Verdict"),
81
+ gr.Textbox(label="AI Probability Score")
82
+ ],
83
+ title="SourceCodeAuthorCheck SLM (10M)",
84
+ description="Analyze Python snippets to determine if they were written by a human or generated by an AI model.",
85
+ examples=[
86
+ ["def calculate_tax(gross_salary, deduction):\n return gross_salary - deduction"],
87
+ ["def process_data_stream_0(data_input: list[dict], strict_validation: bool = True) -> dict:\n if not data_input:\n return {'status': 'error', 'message': 'Empty stream'}\n processed_results = []\n for idx, item in enumerate(data_input):\n transformed = {k: str(v).strip().lower() for k, v in item.items()}\n transformed['_internal_id'] = f'gen_id_0_{idx}'\n processed_results.append(transformed)\n return {'status': 'success', 'data': processed_results}"]
88
+ ]
89
+ )
90
+
91
+ if __name__ == "__main__":
92
+ demo.launch()