upwork-grant-project / worker_script.py
ravi2814's picture
Upload 3 files
e512f47 verified
Raw
History Blame Contribute Delete
12.6 kB
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)