| """Verify access to existing kronlingua / FisherKronecker assets per Task 0.2. | |
| Run: | |
| python scripts/verify_assets.py [--skip-llm] [--skip-bert] | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import os | |
| import sys | |
| from pathlib import Path | |
| KRONLINGUA = Path(os.environ.get("MFFLORA_KRONLINGUA_ROOT", "external/kronlingua")) | |
| FACTORS_ROOT = KRONLINGUA / "factors" / "e1_step2_llama31_base_12lang_gate" | |
| EXPECTED_LANGUAGES = {"ar", "de", "en", "es", "fr", "hi", "ru", "sw", "tr", "ur", "vi", "zh"} | |
| EXPECTED_LAYERS = (0, 1, 14, 15, 30, 31) | |
| def _import_kronlingua(): | |
| if str(KRONLINGUA) not in sys.path: | |
| sys.path.insert(0, str(KRONLINGUA)) | |
| from mff.factors import load_kron_factors # type: ignore | |
| return load_kron_factors | |
| def check_factor_library() -> dict: | |
| load_kron_factors = _import_kronlingua() | |
| sample = FACTORS_ROOT / "en" / "model__layers__0__mlp__gate_proj.safetensors" | |
| if not sample.exists(): | |
| raise FileNotFoundError("Configured factor library is missing the sample factor") | |
| factors = load_kron_factors(str(sample)) | |
| a_shape = tuple(factors.A.shape) | |
| b_shape = tuple(factors.B.shape) | |
| if a_shape != (14336, 14336): | |
| raise AssertionError(f"A shape {a_shape} != (14336, 14336)") | |
| if b_shape != (4096, 4096): | |
| raise AssertionError(f"B shape {b_shape} != (4096, 4096)") | |
| languages = sorted(p.name for p in FACTORS_ROOT.iterdir() if p.is_dir() and len(p.name) <= 3) | |
| missing_langs = EXPECTED_LANGUAGES - set(languages) | |
| if missing_langs: | |
| raise AssertionError(f"Missing languages: {sorted(missing_langs)}") | |
| per_lang_layers: dict[str, list[int]] = {} | |
| for lang in sorted(EXPECTED_LANGUAGES): | |
| files = list((FACTORS_ROOT / lang).glob("model__layers__*__mlp__gate_proj.safetensors")) | |
| layer_ids = sorted(int(f.stem.split("__")[2]) for f in files) | |
| per_lang_layers[lang] = layer_ids | |
| missing = set(EXPECTED_LAYERS) - set(layer_ids) | |
| if missing: | |
| raise AssertionError(f"Lang {lang} missing layers {sorted(missing)}") | |
| return { | |
| "factors_root": "configured_external_factor_root", | |
| "sample_A_shape": a_shape, | |
| "sample_B_shape": b_shape, | |
| "sample_A_dtype": str(factors.A.dtype), | |
| "sample_B_dtype": str(factors.B.dtype), | |
| "languages": languages, | |
| "expected_layers": list(EXPECTED_LAYERS), | |
| "per_language_layers": per_lang_layers, | |
| } | |
| def check_bert() -> dict: | |
| import torch | |
| from transformers import AutoModelForMaskedLM, AutoTokenizer | |
| model_id = "bert-base-uncased" | |
| tok = AutoTokenizer.from_pretrained(model_id) | |
| model = AutoModelForMaskedLM.from_pretrained(model_id) | |
| inputs = tok("hello world", return_tensors="pt") | |
| with torch.no_grad(): | |
| out = model(**inputs) | |
| return {"bert_logits_shape": tuple(out.logits.shape)} | |
| def check_llama(model_path: str) -> dict: | |
| import torch | |
| from transformers import AutoModelForCausalLM, AutoTokenizer | |
| if not Path(model_path).exists(): | |
| raise FileNotFoundError("Configured Llama model is missing") | |
| tok = AutoTokenizer.from_pretrained(model_path) | |
| if tok.pad_token is None: | |
| tok.pad_token = tok.eos_token | |
| model = AutoModelForCausalLM.from_pretrained( | |
| model_path, | |
| torch_dtype=torch.bfloat16, | |
| device_map="cuda" if torch.cuda.is_available() else "cpu", | |
| ) | |
| inputs = tok("Hello, world.", return_tensors="pt").to(model.device) | |
| with torch.no_grad(): | |
| out = model(**inputs) | |
| return {"llama_logits_shape": tuple(out.logits.shape), "device": str(model.device)} | |
| def main() -> None: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--skip-bert", action="store_true") | |
| parser.add_argument("--skip-llm", action="store_true") | |
| parser.add_argument( | |
| "--llama-path", | |
| default=os.environ.get("MFFLORA_MODEL_PATH", "unsloth/Llama-3.1-8B"), | |
| ) | |
| args = parser.parse_args() | |
| report: dict = {} | |
| print("[1/3] checking factor library...", flush=True) | |
| report["factors"] = check_factor_library() | |
| print(" ok:", report["factors"]["sample_A_shape"], report["factors"]["sample_B_shape"]) | |
| if not args.skip_bert: | |
| print("[2/3] loading BERT-base...", flush=True) | |
| report["bert"] = check_bert() | |
| print(" ok:", report["bert"]["bert_logits_shape"]) | |
| else: | |
| print("[2/3] skipping BERT check") | |
| if not args.skip_llm: | |
| print("[3/3] loading Llama-3.1-8B...", flush=True) | |
| report["llama"] = check_llama(args.llama_path) | |
| print(" ok:", report["llama"]["llama_logits_shape"], "on", report["llama"]["device"]) | |
| else: | |
| print("[3/3] skipping Llama check") | |
| out = Path("outputs/verify_assets.json") | |
| out.parent.mkdir(parents=True, exist_ok=True) | |
| out.write_text(json.dumps(report, indent=2)) | |
| print(f"wrote {out}") | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 4.93 kB
- Xet hash:
- ec26856aed39045f3ea8cb061f97a72da5ca53f4b518c696b3031a90dbd77495
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.