tomaszki/PythonFileCompressor
1
1import streamlit as st2from io import StringIO 3from transformers import AutoTokenizer, AutoModelForCausalLM4from torch.nn import functional as F5import torch6import numpy as np7 8import numpyAc9 10st.set_page_config(layout="wide")11device = 'cuda' if torch.cuda.is_available() else 'cpu'12 13@st.cache_resource14def load_model():15 return AutoModelForCausalLM.from_pretrained(16 "PY007/TinyLlama-1.1B-python-v0.1",17 ).to(device)18 19@st.cache_resource20def load_tokenizer():21 return AutoTokenizer.from_pretrained("PY007/TinyLlama-1.1B-python-v0.1")22 23model = load_model()24tokenizer = load_tokenizer()25 26st.title('Python file compressor')27encode_col, decode_col = st.columns(2, gap='medium')28 29@st.cache_data30def encode(text):31 bar = st.progress(0.0)32 codec = numpyAc.arithmeticCoding()33 tokenized = tokenizer(text, return_tensors='pt').input_ids.to(device)34 output = list()35 past_key_values = None36 37 # We can't run a single pass over all tokens, because38 # we get inconsistent results then39 length = tokenized.shape[1]40 for i in range(length):41 bar.progress(min(((i + 1) + (i + 1) ** 2 / 1000) / (length + length ** 2 // 1000), 1.0))42 with torch.no_grad():43 output_ = model(44 input_ids=tokenized[:, i:i + 1],45 use_cache=True,46 past_key_values=past_key_values47 )48 past_key_values = output_.past_key_values49 logits = output_.logits[0, -1:, :]50 output.append(logits)51 output = torch.cat(output, dim=0)52 output = F.softmax(output, dim=-1)53 tokenized = torch.cat((tokenized.squeeze()[1:], torch.tensor([2], device=device))) # Add EOS54 tokenized = tokenized.type(torch.int16).cpu().numpy()55 byte_stream, _ = codec.encode(output.cpu(), tokenized)56 return byte_stream57 58@st.cache_data59def decode(byte_stream):60 # Unfortunately progressbar for decoding isn't possible/is hard61 decodec = numpyAc.arithmeticDeCoding(byte_stream, 32_000)62 input_ids = [1]63 past_key_values = None64 65 while input_ids[-1] != 2:66 with torch.no_grad():67 output = model(68 input_ids=torch.tensor([input_ids[-1:]], device=device),69 use_cache=True,70 past_key_values=past_key_values71 )72 past_key_values = output.past_key_values73 logits = output.logits[0, -1:, :]74 logits = F.softmax(logits, dim=-1).cpu()75 next_token = decodec.decode(logits)76 input_ids.append(next_token)77 return input_ids78 79with encode_col:80 st.header('Convert your python file to binary.')81 python_file = st.file_uploader("Upload your python file here. I recommend files up to 10-20 lines, so it doesn't take too long.")82 if python_file is not None:83 stringio = StringIO(python_file.getvalue().decode("utf-8"))84 code = stringio.read()85 bytes_stream = encode(code)86 bin_filename = f'{python_file.name.split(".")[0]}.bin'87 st.download_button('Download binary file', bytes_stream, bin_filename)88 89with decode_col:90 st.header('Convert your binary file to python')91 binary_file = st.file_uploader('Upload your binary file here')92 if binary_file is not None:93 tokens = decode(binary_file.read())94 decompressed = tokenizer.decode(tokens, skip_special_tokens=True)95 py_filename = f'{binary_file.name.split(".")[0]}.py'96 st.download_button('Download python file', decompressed, py_filename)97 st.code(decompressed)