Spaces:
Sleeping
Sleeping
| import os | |
| import sys | |
| import time | |
| import requests | |
| import json | |
| import shutil | |
| from huggingface_hub import HfApi, hf_hub_download | |
| from datetime import datetime | |
| # --- CONFIGURATION --- | |
| REPO_ID = "ravi2814/grant-data-storage" | |
| STATUS_FILE = "status.json" | |
| def get_token(): | |
| try: | |
| from kaggle_secrets import UserSecretsClient | |
| token = UserSecretsClient().get_secret("HF_TOKEN") | |
| if token: return token | |
| except: pass | |
| token = os.getenv("HF_TOKEN") | |
| if token and "PLACEHOLDER" not in token: return token | |
| return "HF_TOKEN_PLACEHOLDER" | |
| # --- STATUS MANAGER (Heartbeat) --- | |
| class StatusManager: | |
| def __init__(self, repo_id, token): | |
| self.repo_id = repo_id | |
| self.api = HfApi(token=token) | |
| self.last_upload = 0 | |
| def update(self, state, phase, logs, progress): | |
| """Uploads status.json to HF so the UI can read it remotely.""" | |
| try: | |
| # Rate limit: Update at most every 2 seconds | |
| if time.time() - self.last_upload < 2 and state == "running": | |
| return | |
| data = { | |
| "kaggle_state": state, | |
| "phase": phase, | |
| "logs": logs, | |
| "progress": progress, | |
| "timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S") | |
| } | |
| # Write locally | |
| with open(STATUS_FILE, "w") as f: json.dump(data, f) | |
| # Upload to Dataset | |
| self.api.upload_file( | |
| path_or_fileobj=STATUS_FILE, | |
| path_in_repo=STATUS_FILE, | |
| repo_id=self.repo_id, | |
| repo_type="dataset" | |
| ) | |
| self.last_upload = time.time() | |
| except: pass | |
| # --- LOGGER --- | |
| class Logger: | |
| def __init__(self, repo_id, token, status_mgr): | |
| self.repo_id = repo_id | |
| self.api = HfApi(token=token) | |
| self.status_mgr = status_mgr | |
| self.log_file = "live_log.txt" | |
| self.buffer = [] | |
| self.last_upload = 0 | |
| # Init file | |
| with open(self.log_file, "w") as f: | |
| f.write(f"Worker Started at {datetime.now().strftime('%H:%M:%S')}\n") | |
| self.flush() | |
| def log(self, message): | |
| print(message, flush=True) | |
| timestamp = datetime.now().strftime('%H:%M:%S') | |
| log_entry = f"[{timestamp}] {message}" | |
| self.buffer.append(log_entry) | |
| # Update Status JSON immediately with log snippet | |
| self.status_mgr.update("running", "Connected to Kaggle (Safe to close)", log_entry, 95) | |
| if time.time() - self.last_upload > 3: | |
| self.flush() | |
| def flush(self): | |
| try: | |
| with open(self.log_file, "w") as f: f.write("\n".join(self.buffer)) | |
| self.api.upload_file(path_or_fileobj=self.log_file, path_in_repo=self.log_file, repo_id=self.repo_id, repo_type="dataset") | |
| self.last_upload = time.time() | |
| except: pass | |
| if __name__ == "__main__": | |
| HF_TOKEN = get_token() | |
| if "PLACEHOLDER" in HF_TOKEN: sys.exit(1) | |
| # 1. Start Status Manager | |
| status_mgr = StatusManager(REPO_ID, HF_TOKEN) | |
| status_mgr.update("running", "Initializing Environment...", "Starting...", 70) | |
| try: | |
| logger = Logger(REPO_ID, HF_TOKEN, status_mgr) | |
| logger.log("Environment Ready.") | |
| except: sys.exit(1) | |
| # 2. Install Libraries | |
| try: | |
| import sentence_transformers | |
| import pyarrow | |
| except ImportError: | |
| logger.log("Installing libraries...") | |
| os.system("pip install sentence-transformers pyarrow numpy pandas huggingface_hub torch") | |
| # 3. Imports (Fixed Torch Error) | |
| import pandas as pd | |
| import hashlib | |
| import re | |
| import numpy as np | |
| import pyarrow.parquet as pq | |
| import pyarrow as pa | |
| # Explicitly import torch with error handling | |
| try: | |
| import torch | |
| except ImportError: | |
| logger.log("Re-installing torch...") | |
| os.system("pip install torch") | |
| import torch | |
| from sentence_transformers import SentenceTransformer | |
| class Updater: | |
| def __init__(self, logger_instance): | |
| self.logger = logger_instance | |
| self.data_url = "https://open.canada.ca/data/dataset/432527ab-7aac-45b5-81d6-7597107a7013/resource/1d15a62f-5656-49ad-8c88-f40ce689d831/download/grants.csv" | |
| self.api_url = "https://open.canada.ca/data/api/3/action/package_show?id=432527ab-7aac-45b5-81d6-7597107a7013" | |
| self.local_parquet = "grant data.parquet" | |
| self.temp_csv = "temp_data.csv" | |
| self.api_metadata = {} | |
| device = 'cuda' if torch.cuda.is_available() else 'cpu' | |
| self.logger.log(f"Loading AI Model on {device.upper()}...") | |
| self.model = SentenceTransformer('intfloat/multilingual-e5-large', device=device) | |
| def create_hash(self, text): | |
| if not isinstance(text, str): text = "" | |
| clean = re.sub(r'\s+', '', text).lower() | |
| return hashlib.md5(clean.encode('utf-8')).hexdigest() | |
| def format_record(self, record): | |
| def val(k): return str(record.get(k, "") or "").strip() | |
| mapping = [ | |
| ("Reference number", "ref_number"), ("amendment number", "amendment_number"), | |
| ("agreement type", "agreement_type"), ("recipient legal name", "recipient_legal_name"), | |
| ("recipient country", "recipient_country"), ("recipient province", "recipient_province"), | |
| ("recipient city", "recipient_city"), ("recipient postal code", "recipient_postal_code"), | |
| ("program name english", "prog_name_en"), ("program name french", "prog_name_fr"), | |
| ("program purpose english", "prog_purpose_en"), ("program purpose french", "prog_purpose_fr"), | |
| ("agreement title english", "agreement_title_en"), ("agreement title french", "agreement_title_fr"), | |
| ("agreement number", "agreement_number"), ("agreement value", "agreement_value"), | |
| ("agreement start date", "agreement_start_date"), ("agreement end date", "agreement_end_date"), | |
| ("description english", "description_en"), ("description french", "description_fr"), | |
| ("expected results english", "expected_results_en"), ("expected results french", "expected_results_fr"), | |
| ("additional information english", "additional_information_en"), ("additional information french", "additional_information_fr"), | |
| ("owner organization", "owner_org"), ("owner organization title", "owner_org_title"), | |
| ] | |
| return ", ".join([f"{d} : {val(k)}" for d, k in mapping]) | |
| def download_current_database(self): | |
| self.logger.log("Downloading current Database...") | |
| try: | |
| hf_hub_download(repo_id=REPO_ID, filename="grant data.parquet", repo_type="dataset", token=HF_TOKEN, local_dir=".", force_download=True) | |
| hf_hub_download(repo_id=REPO_ID, filename="embeddings.npy", repo_type="dataset", token=HF_TOKEN, local_dir=".", force_download=True) | |
| except: | |
| self.logger.log("No existing database found. Starting fresh.") | |
| def download_government_data(self): | |
| self.logger.log("Downloading fresh data from Canada.ca...") | |
| try: | |
| r = requests.get(self.data_url, stream=True) | |
| with open(self.temp_csv, 'wb') as f: | |
| for chunk in r.iter_content(chunk_size=1048576): f.write(chunk) | |
| r_api = requests.get(self.api_url).json()['result'] | |
| res = next((r for r in r_api['resources'] if r['id'] == '1d15a62f-5656-49ad-8c88-f40ce689d831'), {}) | |
| self.api_metadata = { | |
| "metadata_modified": r_api.get('metadata_modified'), | |
| "resource_modified": res.get('last_modified'), | |
| "checked_at": datetime.now().isoformat() | |
| } | |
| return True | |
| except Exception as e: | |
| self.logger.log(f"Download Error: {e}") | |
| return False | |
| def synchronize_data(self): | |
| if not os.path.exists(self.temp_csv): return False | |
| self.logger.log("Scanning Government CSV...") | |
| status_mgr.update("running", "Scanning CSV...", "Reading...", 60) | |
| valid_hashes = set() | |
| new_rows_buffer = [] | |
| for chunk in pd.read_csv(self.temp_csv, chunksize=50000, low_memory=False): | |
| for _, row in chunk.iterrows(): | |
| fmt = self.format_record(row.to_dict()) | |
| h = self.create_hash(fmt) | |
| valid_hashes.add(h) | |
| new_rows_buffer.append({'data': fmt, 'hash': h}) | |
| self.logger.log(f"Total Official Records: {len(valid_hashes)}") | |
| if os.path.exists(self.local_parquet) and os.path.exists("embeddings.npy"): | |
| df = pd.read_parquet(self.local_parquet) | |
| emb = np.load("embeddings.npy") | |
| keep_indices = [] | |
| local_hashes = set() | |
| for idx, row in df.iterrows(): | |
| h = self.create_hash(row['data']) | |
| if h in valid_hashes: | |
| keep_indices.append(idx) | |
| local_hashes.add(h) | |
| if len(keep_indices) < len(df): | |
| self.logger.log(f"Removing {len(df) - len(keep_indices)} obsolete records...") | |
| df = df.loc[keep_indices].reset_index(drop=True) | |
| emb = emb[keep_indices] | |
| else: | |
| df = pd.DataFrame(columns=['data']) | |
| emb = np.empty((0, 1024)) | |
| local_hashes = set() | |
| rows_to_add = [] | |
| for item in new_rows_buffer: | |
| if item['hash'] not in local_hashes: | |
| rows_to_add.append(item['data']) | |
| local_hashes.add(item['hash']) | |
| if rows_to_add: | |
| self.logger.log(f"Found {len(rows_to_add)} NEW records. Processing...") | |
| status_mgr.update("running", "Embedding New Data...", f"Processing {len(rows_to_add)} rows", 80) | |
| new_vectors = [] | |
| for i in range(0, len(rows_to_add), 2000): | |
| batch = rows_to_add[i : i + 2000] | |
| batch_fmt = ["passage: " + t for t in batch] | |
| v = self.model.encode(batch_fmt, convert_to_numpy=True, normalize_embeddings=True) | |
| new_vectors.append(v) | |
| new_emb_stack = np.vstack(new_vectors) | |
| new_df = pd.DataFrame({'data': rows_to_add}) | |
| df = pd.concat([df, new_df], ignore_index=True) | |
| if len(emb) > 0: emb = np.vstack([emb, new_emb_stack]) | |
| else: emb = new_emb_stack | |
| self.logger.log(f"Database Updated. Total: {len(df)}") | |
| else: | |
| self.logger.log("No new records to add.") | |
| self.logger.log("Saving files...") | |
| table = pa.Table.from_pandas(df) | |
| pq.write_table(table, self.local_parquet) | |
| np.save("embeddings.npy", emb) | |
| with open("last_metadata.json", "w") as f: | |
| json.dump(self.api_metadata, f) | |
| return True | |
| try: | |
| updater = Updater(logger) | |
| updater.download_current_database() | |
| if updater.download_government_data(): | |
| updater.synchronize_data() | |
| logger.log("Uploading results...") | |
| status_mgr.update("running", "Uploading Final Data...", "Finalizing...", 95) | |
| api = HfApi(token=HF_TOKEN) | |
| os.makedirs("staging", exist_ok=True) | |
| files = ["grant data.parquet", "embeddings.npy", "last_metadata.json", "live_log.txt", STATUS_FILE] | |
| for f in files: | |
| if os.path.exists(f): shutil.copy(f, f"staging/{f}") | |
| api.upload_folder(folder_path="staging", repo_id=REPO_ID, repo_type="dataset", commit_message="Batch Update") | |
| # FINAL SUCCESS | |
| status_mgr.update("finished", "Sync Complete", "All data uploaded successfully.", 100) | |
| logger.log("SUCCESS: Process Complete.") | |
| except Exception as e: | |
| logger.log(f"CRITICAL ERROR: {e}") | |
| status_mgr.update("error", "Worker Failed", str(e), 0) | |
| sys.exit(1) |