BanBTP-2.0-V10 / chat.py
aaro765's picture
Upload 10 files
32096b9 verified
Raw
History Blame Contribute Delete
15.4 kB
import os
import sys
import subprocess
import venv
import json
import math
import argparse
import shutil
import csv
import xml.etree.ElementTree as ET
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
TEMP_DIR = os.path.join(SCRIPT_DIR, ".temp")
MODEL_DIR = os.path.join(SCRIPT_DIR, "banbtp-final")
DATASET_DIR = os.path.join(SCRIPT_DIR, "dataset")
BANNER_STATE_FILE = os.path.join(MODEL_DIR, "banner_state.pt")
os.makedirs(TEMP_DIR, exist_ok=True)
os.makedirs(DATASET_DIR, exist_ok=True)
os.environ["HF_HOME"] = os.path.join(TEMP_DIR, "hf_home")
os.environ["TOKENIZERS_PARALLELISM"] = "false"
def bootstrap_venv():
venv_dir = os.path.join(TEMP_DIR, "venv")
python_exe = os.path.join(venv_dir, "Scripts" if os.name == "nt" else "bin", "python")
if os.path.realpath(sys.executable) != os.path.realpath(python_exe):
if not os.path.exists(python_exe):
print(">>> Creating isolated virtual environment...")
venv.create(venv_dir, with_pip=True)
subprocess.check_call([python_exe, "-m", "pip", "install", "--upgrade", "pip", "-q"], env=os.environ)
subprocess.check_call([
python_exe, "-m", "pip", "install",
"torch", "transformers>=4.38.0", "safetensors", "sentencepiece", "tqdm",
"pandas", "pyarrow", "accelerate", "-q"
], env=os.environ)
os.execv(python_exe, [python_exe] + sys.argv)
bootstrap_venv()
import torch
import torch.nn.functional as F
from transformers import AutoModelForCausalLM, AutoTokenizer
from tqdm import tqdm
import pandas as pd
import pyarrow.parquet as pq
import pyarrow as pa
# ==============================================================================
# Banner Engine (Extra Parameters / Non-Parametric Memory)
# ==============================================================================
class BannerEngine:
def __init__(self, model, tokenizer, device):
self.model = model
self.tokenizer = tokenizer
self.device = device
self.markov = {}
self.rag_keys = []
self.rag_texts = []
self.gene_pool = torch.zeros((1000, model.config.hidden_size))
self.fitness = torch.zeros(1000)
self.generation = 0
def extract_text(self, file_path):
ext = os.path.splitext(file_path)[1].lower()
try:
if ext == '.txt':
with open(file_path, 'r', encoding='utf-8', errors='ignore') as f: return f.read()
elif ext == '.json':
with open(file_path, 'r', encoding='utf-8', errors='ignore') as f: return json.dumps(json.load(f))
elif ext == '.jsonl':
with open(file_path, 'r', encoding='utf-8', errors='ignore') as f: return '\n'.join([json.dumps(json.loads(line)) for line in f if line.strip()])
elif ext == '.csv':
df = pd.read_csv(file_path, on_bad_lines='skip')
return ' '.join(df.astype(str).agg(' '.join, axis=1).tolist())
elif ext == '.tsv':
df = pd.read_csv(file_path, sep='\t', on_bad_lines='skip')
return ' '.join(df.astype(str).agg(' '.join, axis=1).tolist())
elif ext == '.parquet':
df = pq.read_table(file_path).to_pandas()
return ' '.join(df.astype(str).agg(' '.join, axis=1).tolist())
elif ext in ['.arrow', '.feather']:
reader = pa.ipc.RecordBatchFileReader(pa.memory_map(file_path))
df = reader.read_all().to_pandas()
return ' '.join(df.astype(str).agg(' '.join, axis=1).tolist())
elif ext == '.xml':
tree = ET.parse(file_path)
return ''.join(tree.getroot().itertext())
except Exception as e:
print(f" [Warning] Could not parse {os.path.basename(file_path)}: {e}")
return ""
def ingest_text(self, text, chunk_size=256):
if not text.strip(): return
tokens = self.tokenizer.encode(text, add_special_tokens=False)
for i in range(len(tokens)-1):
prev, nxt = tokens[i], tokens[i+1]
if prev not in self.markov: self.markov[prev] = {}
self.markov[prev][nxt] = self.markov[prev].get(nxt, 0) + 1
for i in range(0, len(tokens), chunk_size):
chunk_tokens = tokens[i:i+chunk_size]
if len(chunk_tokens) < 10: continue
chunk_text = self.tokenizer.decode(chunk_tokens, skip_special_tokens=True)
input_ids = torch.tensor([chunk_tokens], dtype=torch.long).to(self.device)
with torch.no_grad():
embeds = self.model.model.embed_tokens(input_ids)
chunk_embed = embeds.mean(dim=1).squeeze(0).cpu()
self.rag_keys.append(chunk_embed)
self.rag_texts.append(chunk_text)
if tokens:
input_ids = torch.tensor([tokens[:256]], dtype=torch.long).to(self.device)
with torch.no_grad():
embeds = self.model.model.embed_tokens(input_ids)
self.adapt_genetic(embeds.mean(dim=1).squeeze(0).cpu(), 0.5)
def adapt_genetic(self, hidden_state, error):
weakest = torch.argmin(self.fitness)
mutation = torch.randn_like(self.gene_pool[weakest]) * error
self.gene_pool[weakest] = hidden_state + mutation
self.fitness[weakest] = 1.0 / (error + 1e-5)
self.generation += 1
def retrieve(self, query_text, top_k=2):
if not self.rag_keys: return [], 0.0
tokens = self.tokenizer.encode(query_text, add_special_tokens=False)[:256]
if not tokens: return [], 0.0
input_ids = torch.tensor([tokens], dtype=torch.long).to(self.device)
with torch.no_grad():
q_embed = self.model.model.embed_tokens(input_ids).mean(dim=1).squeeze(0).cpu()
keys_tensor = torch.stack(self.rag_keys)
sims = F.cosine_similarity(q_embed.unsqueeze(0), keys_tensor, dim=1)
top_sims, top_indices = torch.topk(sims, k=min(top_k, len(self.rag_texts)))
max_sim = top_sims[0].item() if len(top_sims) > 0 else 0.0
retrieved = [self.rag_texts[i] for i in top_indices.tolist()]
return retrieved, max_sim
def save_state(self):
state = {
"markov": self.markov,
"rag_keys": self.rag_keys,
"rag_texts": self.rag_texts,
"gene_pool": self.gene_pool,
"fitness": self.fitness,
"generation": self.generation
}
torch.save(state, BANNER_STATE_FILE)
def load_state(self):
if not os.path.exists(BANNER_STATE_FILE): return False
state = torch.load(BANNER_STATE_FILE, map_location="cpu", weights_only=False)
self.markov = state.get("markov", {})
self.rag_keys = state.get("rag_keys", [])
self.rag_texts = state.get("rag_texts", [])
self.gene_pool = state.get("gene_pool", torch.zeros((1000, self.model.config.hidden_size)))
self.fitness = state.get("fitness", torch.zeros(1000))
self.generation = state.get("generation", 0)
return True
# ==============================================================================
# Finetune Mode
# ==============================================================================
def run_finetune(banner):
print(f">>> Scanning {DATASET_DIR} for dataset files...")
supported_ext = ('.txt', '.json', '.jsonl', '.csv', '.tsv', '.xml', '.parquet', '.arrow', '.feather')
files = []
for root, dirs, filenames in os.walk(DATASET_DIR):
for f in filenames:
if f.lower().endswith(supported_ext):
files.append(os.path.join(root, f))
if not files:
print(">>> No supported files found in dataset folder.")
return
print(f">>> Found {len(files)} files. Ingesting into Banner extra parameters...")
for fpath in tqdm(files, desc="Processing Files", unit="file"):
text = banner.extract_text(fpath)
banner.ingest_text(text)
try:
os.remove(fpath)
except Exception as e:
print(f" [Warning] Could not delete {fpath}: {e}")
for root, dirs, files in os.walk(DATASET_DIR, topdown=False):
for name in dirs:
dir_path = os.path.join(root, name)
if dir_path != DATASET_DIR:
try:
os.rmdir(dir_path)
except OSError:
pass
banner.save_state()
print(f">>> Finetune complete! Banner state saved to {BANNER_STATE_FILE}")
print(f" Markov transitions: {sum(len(v) for v in banner.markov.values()):,}")
print(f" RAG memories: {len(banner.rag_texts):,}")
print(f" Genetic generation: {banner.generation}")
print(f">>> Dataset folder cleared automatically.")
# ==============================================================================
# Chat Mode
# ==============================================================================
def run_chat(banner, model, tokenizer, device):
temp = 0.7
max_tokens = 64
auto_temp = True
print("=" * 50)
print(" banbtp2.0v10 chat + Banner Engine")
print(" type /help for commands")
print("=" * 50)
while True:
try:
user_input = input("\nyou> ").strip()
except (EOFError, KeyboardInterrupt):
print("\n>>> Saving banner state...")
banner.save_state()
break
if not user_input: continue
if user_input.startswith("/"):
cmd_parts = user_input.split()
cmd = cmd_parts[0].lower()
if cmd in ("/quit", "/exit", "/q"):
print(">>> Saving banner state...")
banner.save_state()
break
elif cmd == "/help":
print(" /temp <0.1-2.0> set manual temperature")
print(" /auto toggle auto-temperature (RAG adaptive)")
print(" /tokens <n> set max new tokens")
print(" /stats show banner memory stats")
print(" /save save banner state now")
print(" /quit save and exit")
elif cmd == "/temp":
if len(cmd_parts) > 1:
try:
temp = max(0.1, min(2.0, float(cmd_parts[1])))
auto_temp = False
print(f" manual temperature = {temp}")
except ValueError: print(" usage: /temp <number>")
else: print(f" temperature = {temp}")
elif cmd == "/auto":
auto_temp = not auto_temp
print(f" auto-temperature {'ON' if auto_temp else 'OFF'}")
elif cmd == "/tokens":
if len(cmd_parts) > 1:
try:
max_tokens = max(1, int(cmd_parts[1]))
print(f" max_new_tokens = {max_tokens}")
except ValueError: print(" usage: /tokens <number>")
else: print(f" max_new_tokens = {max_tokens}")
elif cmd == "/save":
banner.save_state()
print(" banner state saved.")
elif cmd == "/stats":
print(f" markov transitions: {sum(len(v) for v in banner.markov.values()):,}")
print(f" rag memories: {len(banner.rag_texts):,}")
print(f" genetic generation: {banner.generation}")
else:
print(f" unknown command: {cmd}. type /help")
continue
retrieved, max_sim = banner.retrieve(user_input, top_k=2)
if retrieved and max_sim > 0.5:
context = "\n".join(retrieved)
full_prompt = f"{context}\n{user_input}"
current_temp = max(0.2, temp * 0.5) if auto_temp else temp
else:
full_prompt = user_input
current_temp = temp * 1.2 if auto_temp else temp
input_ids = tokenizer(full_prompt, return_tensors="pt", truncation=True, max_length=512)["input_ids"].to(device)
if input_ids.shape[1] == 0:
print("model> [Empty prompt]")
continue
generated = []
with torch.no_grad():
for _ in tqdm(range(max_tokens), desc="Thinking", bar_format='{l_bar}{bar}| {n_fmt}/{total_fmt}', leave=False):
outputs = model(input_ids)
next_token_logits = outputs.logits[:, -1, :] / current_temp
if torch.isnan(next_token_logits).any():
next_token_logits = torch.nan_to_num(next_token_logits, nan=0.0)
probs = F.softmax(next_token_logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
generated.append(next_token.item())
if next_token.item() == tokenizer.eos_token_id:
break
input_ids = torch.cat([input_ids, next_token], dim=-1)
if input_ids.shape[1] > 1024:
input_ids = input_ids[:, -1024:]
response = tokenizer.decode(generated, skip_special_tokens=True)
out_tokens = tokenizer.encode(response, add_special_tokens=False)
for i in range(len(out_tokens)-1):
prev, nxt = out_tokens[i], out_tokens[i+1]
if prev not in banner.markov: banner.markov[prev] = {}
banner.markov[prev][nxt] = banner.markov[prev].get(nxt, 0) + 1
if out_tokens:
input_ids_embed = torch.tensor([out_tokens[:256]], dtype=torch.long).to(device)
with torch.no_grad():
embeds = model.model.embed_tokens(input_ids_embed)
banner.adapt_genetic(embeds.mean(dim=1).squeeze(0).cpu(), 0.1)
print(f"\nmodel> {response}")
# ==============================================================================
# Main
# ==============================================================================
def main():
parser = argparse.ArgumentParser(description="banbtp2.0v10 chat interface")
parser.add_argument("--chat", action="store_true", help="chat mode (default)")
parser.add_argument("--finetune", action="store_true", help="ingest dataset folder into banner memory")
args = parser.parse_args()
if not os.path.exists(MODEL_DIR):
print(f"ERROR: Model directory not found: {MODEL_DIR}")
sys.exit(1)
print(f">>> Loading banbtp2.0v10 from: {MODEL_DIR}")
tokenizer = AutoTokenizer.from_pretrained(MODEL_DIR)
# FIX: Removed device_map="auto" to prevent accelerate crash.
# Added trust_remote_code=True so it loads the custom arch silently.
model = AutoModelForCausalLM.from_pretrained(
MODEL_DIR,
dtype=torch.float32,
trust_remote_code=True
)
model.eval()
device = next(model.parameters()).device
banner = BannerEngine(model, tokenizer, device)
if banner.load_state():
print(">>> Loaded existing Banner state.")
if args.finetune:
run_finetune(banner)
else:
run_chat(banner, model, tokenizer, device)
if __name__ == "__main__":
main()