JerMa88's picture
Upload folder using huggingface_hub
3ce19a2 verified
Raw
History Blame Contribute Delete
2.61 kB
import os
import json
from huggingface_hub import HfApi
def get_param_count():
# The default log path from train_cifar10.sh and train.py
log_file = "./runs/cifar10_rtm/cifar10/log.jsonl"
if not os.path.exists(log_file):
return "unknown"
with open(log_file, "r") as f:
for line in f:
try:
data = json.loads(line)
if "message" in data and "Number of parameters in IMLE:" in data["message"]:
parts = data["message"].split("Number of parameters in IMLE:")
if len(parts) > 1:
params_str = parts[1].strip()
return params_str
except:
pass
return "unknown"
def main():
api = HfApi()
# 1. Get parameter count and format repo name
param_count = get_param_count()
repo_id = f"JerMa88/rtm-latent-refinement-cifar10-{param_count}params"
print(f"Creating repository: {repo_id}")
# Create the repo (it will use the implicitly logged-in user)
try:
api.create_repo(repo_id, exist_ok=True)
except Exception as e:
print(f"Error creating repo: {e}")
# Upload training checkpoints and logs
run_dir = "./runs/cifar10_rtm/cifar10"
if os.path.exists(run_dir):
print(f"Uploading run directory {run_dir}...")
api.upload_folder(
folder_path=run_dir,
repo_id=repo_id,
path_in_repo="runs/cifar10",
)
else:
print(f"Run directory {run_dir} not found! Check if training finished successfully.")
# Upload current repository code
print("Uploading repository code...")
api.upload_folder(
folder_path=".",
repo_id=repo_id,
path_in_repo="code",
ignore_patterns=["runs/*", "datasets/*", "venv/*", ".git/*"]
)
# Upload README to Hugging Face
readme_content = f"""
# RTM Latent Refinement (Fast Variant)
This model was trained with a reduced architecture ({param_count} parameters) compared to the original paper's baseline to speed up training.
- **Dataset**: CIFAR-10
- **W&B Logs**: Check the W&B project `rtm-latent-refinement` under the account `JerMa88`.
- **Source Code**: Included in the `code/` directory.
"""
with open("HF_README.md", "w") as f:
f.write(readme_content)
api.upload_file(
path_or_fileobj="HF_README.md",
path_in_repo="README.md",
repo_id=repo_id,
)
print("Upload complete! You can view the model on Hugging Face.")
if __name__ == "__main__":
main()