Detector / src /streamlit_app.py
InfiniteTsukuyomi's picture
Update src/streamlit_app.py
b5e3324 verified
Raw History Blame Contribute Delete
2.96 kB
import streamlit as st
import torch
from transformers import AutoTokenizer, AutoModelForMaskedLM
import numpy as np
from sklearn.linear_model import LogisticRegression
# Streamlit UI
st.title("Code Analyzer: Tokenization, Probabilities, and Logistic Regression")
# Initialize CodeBERT model and tokenizer with caching to reduce memory usage
@st.cache_resource
def load_model():
tokenizer = AutoTokenizer.from_pretrained("microsoft/codebert-base-mlm")
model = AutoModelForMaskedLM.from_pretrained("microsoft/codebert-base-mlm")
# Optional: Quantize model to reduce memory footprint
model = torch.quantization.quantize_dynamic(model, {torch.nn.Linear}, dtype=torch.qint8)
return tokenizer, model
tokenizer, model = load_model()
# Function to split tokens with 'Ġ' prefix
def split_tokens_with_g(token_list):
new_tokens = []
for token in token_list:
if token.startswith('Ġ') and len(token) > 1:
new_tokens.append('Ġ')
new_tokens.append(token[1:])
else:
new_tokens.append(token)
return new_tokens
# Function to estimate token probabilities
def estimate_probabilities_bidirectional_strategy(code_snippet):
tokens = tokenizer.tokenize(code_snippet)
tokens = split_tokens_with_g(tokens)
token_ids = [tokenizer.convert_tokens_to_ids([t])[0] for t in tokens]
max_seq_length = 511 # 512 - 1 for [CLS]
if len(token_ids) > max_seq_length:
token_ids = token_ids[:max_seq_length]
tokens = tokens[:max_seq_length]
probs = []
for i, token_id in enumerate(token_ids):
input_sequence = token_ids.copy()
input_sequence[i] = tokenizer.mask_token_id
input_ids = [tokenizer.cls_token_id] + input_sequence + [tokenizer.pad_token_id] * (max_seq_length - len(input_sequence))
input_ids = input_ids[:512]
input_tensor = torch.tensor([input_ids])
with torch.no_grad():
outputs = model(input_tensor)
logits = outputs.logits[0, i + 1] # Logits for <mask> position (i + 1 due to [CLS])
probs_token = torch.softmax(logits, dim=-1)[token_id].item()
probs.append(probs_token)
return tokens, probs
# User input: File upload or text input
uploaded_file = st.file_uploader("Upload a code file", type=["py", "java", "txt"])
code_input = st.text_area("Or enter code here", height=200)
if uploaded_file or code_input:
# Read code
code = uploaded_file.read().decode("utf-8") if uploaded_file else code_input
st.write("### Code Input")
st.code(code, language="python" if code.endswith(".py") else "java")
# Step 1: Tokenize and estimate probabilities
try:
tokens, probs = estimate_probabilities_bidirectional_strategy(code)
st.write("### Tokens")
st.write(tokens) # Limit output for readability
st.write("### Token Probabilities")
st.write(probs) # Limit output for readability
except Exception as e:
st.error(f"Error processing code: {str(e)}")
st.write("Upload a code file or enter code to analyze.")