File size: 2,680 Bytes
a2ffd07 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 | from huggingface_hub import login, upload_file, hf_hub_download, delete_file
import os
from dotenv import load_dotenv
def manage_model(action="push"):
load_dotenv()
hf_key = os.getenv('HUGGING_FACE_API_KEY')
login(hf_key)
assert action in ["push", "pull", "delete"]
model_name = "hallucination"
user_name = "ToiTenBao"
repo_id = f"{user_name}/{model_name}"
folder_path = "./cc3m_checkpoints"
file_name = "topk_32.0_32_vision_model.encoder.layers.{layer}.hook_resid_post_0.001_256_0.03125_42.ckpt"
# file_name = "topk_32.0_32_text_decoder.bert.encoder.layer.{layer}.crossattention.self.hook_resid_pre_0.001_256_0.03125_42.ckpt"
# folder_path = "./vis_file"
# file_name = "sae_graph_75719_nonhal.html"
# Removed leading '/'
max_layers = 12
os.makedirs(folder_path, exist_ok=True)
for layer in range(9, max_layers):
base_name = file_name.format(layer=layer)
remote_path = f"cc3m_checkpoints/{base_name}"
local_path = os.path.join(folder_path, base_name)
if action.lower() == "push":
upload_file(
path_or_fileobj=local_path,
path_in_repo=remote_path,
repo_id=repo_id,
)
elif action.lower() == "pull":
local_dir = os.path.dirname(folder_path) if folder_path != "." else "."
hf_hub_download(
repo_id=repo_id,
filename=remote_path,
local_dir=local_dir,
)
elif action.lower() == "delete":
try:
delete_file(
path_in_repo=remote_path,
repo_id=repo_id,
)
except Exception as e:
print(f"Error deleting {remote_path}: {e}")
continue
else:
raise ValueError("Action must be 'push', 'pull', or 'delete'")
def delete_single_file(file_layer):
load_dotenv()
hf_key = os.getenv('HUGGING_FACE_API_KEY')
login(hf_key)
model_name = "hallucination"
user_name = "ToiTenBao"
repo_id = f"{user_name}/{model_name}"
file_name = "topk_16.0_32_text_decoder.bert.encoder.layer.{layer}.attention.self.hook_resid_pre_0.001_256_0.0_42.ckpt"
# Removed leading '/'
base_name = file_name.format(layer=file_layer)
remote_path = f"checkpoints/{base_name}"
try:
delete_file(
path_in_repo=remote_path,
repo_id=repo_id,
)
print(f"Deleted {remote_path}")
except Exception as e:
print(f"Error deleting {remote_path}: {e}")
if __name__ == "__main__":
manage_model(action="pull") |