etomoscow/mff_lora / code /scripts /verify_assets.py
etomoscow's picture
download
raw
4.93 kB
"""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.