Text Generation
PEFT
Safetensors
Arabic
arabic
relation-extraction
qlora
bitsandbytes
multiple-choice
conversational
Instructions to use U4RASD/DRU-RE-Yehia with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use U4RASD/DRU-RE-Yehia with PEFT:
from peft import PeftModel from transformers import AutoModelForCausalLM base_model = AutoModelForCausalLM.from_pretrained("Navid-AI/Yehia-7B-preview") model = PeftModel.from_pretrained(base_model, "U4RASD/DRU-RE-Yehia") - Notebooks
- Google Colab
- Kaggle
| #!/usr/bin/env python3 | |
| """Run constrained one-token inference with a staged DRU-RE-Yehia adapter.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import os | |
| from pathlib import Path | |
| from typing import Any, Dict, Iterable, List, Mapping | |
| os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") | |
| os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "0" | |
| import torch | |
| from dotenv import load_dotenv | |
| from huggingface_hub import snapshot_download | |
| from peft import PeftModel | |
| from tqdm.auto import tqdm | |
| from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig | |
| from re_sft_common import ( | |
| OPTION_CODES, | |
| env_int, | |
| env_str, | |
| load_jsonl, | |
| validate_local_hf_revision, | |
| ) | |
| ROOT = Path(__file__).resolve().parent | |
| load_dotenv(ROOT / ".env") | |
| BASE_MODEL_ID = env_str("YEHIA_BASE_MODEL_ID", "Navid-AI/Yehia-7B-preview") | |
| BASE_MODEL_REVISION = env_str( | |
| "YEHIA_BASE_MODEL_REVISION", "b9dda4715eafee7e8090d2c83cfe078d75f4ebb8" | |
| ) | |
| def repository_path(value: str) -> Path: | |
| path = Path(value) | |
| return path if path.is_absolute() else ROOT / path | |
| def require_token() -> str: | |
| token = os.environ.get("HF_TOKENONE", "").strip() | |
| if not token or token.startswith("hf_your_"): | |
| raise RuntimeError("HF_TOKENONE is required because the local Yehia snapshot is absent") | |
| return token | |
| def resolve_base_model() -> Path: | |
| model_dir = repository_path( | |
| env_str("LOCAL_YEHIA_MODEL_DIR", "models/Yehia-7B-preview") | |
| ) | |
| if (model_dir / "config.json").is_file(): | |
| validate_local_hf_revision(model_dir, BASE_MODEL_REVISION) | |
| return model_dir | |
| offline = os.environ.get("HF_HUB_OFFLINE", "").strip().lower() | |
| if offline in {"1", "true", "yes", "on"}: | |
| raise RuntimeError("Yehia is missing locally while HF_HUB_OFFLINE is enabled") | |
| model_dir.mkdir(parents=True, exist_ok=True) | |
| snapshot_download( | |
| repo_id=BASE_MODEL_ID, | |
| revision=BASE_MODEL_REVISION, | |
| local_dir=str(model_dir), | |
| token=require_token(), | |
| ) | |
| if not (model_dir / "config.json").is_file(): | |
| raise RuntimeError(f"Incomplete Yehia snapshot: {model_dir}") | |
| (model_dir / ".dru_hf_revision").write_text(BASE_MODEL_REVISION + "\n", encoding="utf-8") | |
| validate_local_hf_revision(model_dir, BASE_MODEL_REVISION) | |
| return model_dir | |
| def resolve_adapter() -> Path: | |
| configured = repository_path(env_str("ADAPTER_DIR", ".")) | |
| candidates = [configured, ROOT, ROOT / "runs" / "DRU-RE-Yehia" / "best_adapter"] | |
| for candidate in candidates: | |
| if (candidate / "adapter_config.json").is_file() and ( | |
| candidate / "adapter_model.safetensors" | |
| ).is_file(): | |
| return candidate | |
| raise FileNotFoundError( | |
| "No staged adapter found. Run tools/stage_release.py after training completes " | |
| "or set ADAPTER_DIR to a PEFT adapter directory." | |
| ) | |
| def selected_bias(adapter_dir: Path, override: float | None) -> float: | |
| if override is not None: | |
| return override | |
| for path in (adapter_dir / "inference_config.json", ROOT / "inference_config.json"): | |
| if path.is_file(): | |
| return float(json.loads(path.read_text(encoding="utf-8"))["no_relation_logit_bias"]) | |
| print("WARNING: no inference_config.json found; using no-relation bias 0.0") | |
| return 0.0 | |
| def batches(rows: List[Dict[str, Any]], size: int) -> Iterable[List[Dict[str, Any]]]: | |
| for start in range(0, len(rows), size): | |
| yield rows[start : start + size] | |
| def main() -> None: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument( | |
| "--input", | |
| default=env_str("CHOICE_DATASET_DIR", "data/Yehia-RE-SFT") + "/official.jsonl", | |
| ) | |
| parser.add_argument("--output", default="predictions/official_predictions.jsonl") | |
| parser.add_argument("--batch-size", type=int, default=env_int("INFERENCE_BATCH_SIZE", 16)) | |
| parser.add_argument("--limit", type=int, default=0) | |
| parser.add_argument("--no-relation-bias", type=float, default=None) | |
| args = parser.parse_args() | |
| input_path = repository_path(args.input) | |
| output_path = repository_path(args.output) | |
| rows = load_jsonl(input_path) | |
| if args.limit > 0: | |
| rows = rows[: args.limit] | |
| if not rows: | |
| raise RuntimeError(f"No rows found in {input_path}") | |
| base_dir = resolve_base_model() | |
| adapter_dir = resolve_adapter() | |
| bias = selected_bias(adapter_dir, args.no_relation_bias) | |
| tokenizer_source = adapter_dir if (adapter_dir / "tokenizer_config.json").is_file() else base_dir | |
| tokenizer = AutoTokenizer.from_pretrained( | |
| str(tokenizer_source), local_files_only=True, use_fast=True | |
| ) | |
| if not tokenizer.chat_template: | |
| raise RuntimeError("Yehia tokenizer has no native chat template") | |
| if tokenizer.pad_token_id is None: | |
| tokenizer.pad_token = tokenizer.eos_token | |
| tokenizer.padding_side = "right" | |
| code_token_ids: List[int] = [] | |
| for code in OPTION_CODES: | |
| ids = tokenizer.encode(" " + code, add_special_tokens=False) | |
| if len(ids) != 1: | |
| raise RuntimeError(f"Decision code {code!r} is not one token: {ids}") | |
| code_token_ids.append(int(ids[0])) | |
| quantization = BitsAndBytesConfig( | |
| load_in_4bit=True, | |
| bnb_4bit_quant_type="nf4", | |
| bnb_4bit_use_double_quant=True, | |
| bnb_4bit_compute_dtype=torch.bfloat16, | |
| ) | |
| base = AutoModelForCausalLM.from_pretrained( | |
| str(base_dir), | |
| local_files_only=True, | |
| quantization_config=quantization, | |
| torch_dtype=torch.bfloat16, | |
| device_map={"": 0}, | |
| attn_implementation=env_str("ATTENTION_IMPLEMENTATION", "sdpa"), | |
| ) | |
| model = PeftModel.from_pretrained(base, str(adapter_dir), is_trainable=False) | |
| model.eval() | |
| output_path.parent.mkdir(parents=True, exist_ok=True) | |
| with output_path.open("w", encoding="utf-8") as handle, torch.inference_mode(): | |
| for group in tqdm(list(batches(rows, args.batch_size)), desc="predict"): | |
| sequences = [ | |
| list( | |
| tokenizer.apply_chat_template( | |
| row["prompt_messages"], tokenize=True, add_generation_prompt=True | |
| ) | |
| ) | |
| for row in group | |
| ] | |
| lengths = torch.tensor([len(sequence) for sequence in sequences], dtype=torch.long) | |
| max_length = int(lengths.max().item()) | |
| input_ids = torch.full( | |
| (len(group), max_length), tokenizer.pad_token_id, dtype=torch.long | |
| ) | |
| attention_mask = torch.zeros((len(group), max_length), dtype=torch.long) | |
| for index, sequence in enumerate(sequences): | |
| input_ids[index, : len(sequence)] = torch.tensor(sequence, dtype=torch.long) | |
| attention_mask[index, : len(sequence)] = 1 | |
| input_ids = input_ids.to(model.device) | |
| attention_mask = attention_mask.to(model.device) | |
| logits = model(input_ids=input_ids, attention_mask=attention_mask).logits | |
| next_logits = logits[ | |
| torch.arange(len(group), device=logits.device), | |
| lengths.to(logits.device) - 1, | |
| ] | |
| for index, row in enumerate(group): | |
| options = list(row["allowed_options_ar"]) | |
| labels = list(row["allowed_relation_full_labels"]) | |
| codes = list(row["option_codes"]) | |
| if not options or options[-1] != "لا توجد علاقة": | |
| raise RuntimeError(f"Malformed option list for {row.get('id')}") | |
| candidate_ids = torch.tensor( | |
| code_token_ids[: len(options)], device=logits.device | |
| ) | |
| scores = next_logits[index, candidate_ids].float().clone() | |
| scores[-1] += bias | |
| chosen = int(torch.argmax(scores).item()) | |
| record: Dict[str, Any] = { | |
| "id": row.get("id"), | |
| "sentence_id": row.get("sentence_id"), | |
| "triple_id": row.get("triple_id"), | |
| "predicted_option_index": chosen, | |
| "predicted_code": codes[chosen], | |
| "predicted_option_ar": options[chosen], | |
| "predicted_relation_full": labels[chosen], | |
| "no_relation_logit_bias": bias, | |
| } | |
| ontology_ids = row.get("allowed_relation_ontology_ids") | |
| if isinstance(ontology_ids, list) and chosen < len(ontology_ids): | |
| record["predicted_relation_ontology_id"] = ontology_ids[chosen] | |
| handle.write(json.dumps(record, ensure_ascii=False) + "\n") | |
| print(json.dumps({"rows": len(rows), "output": str(output_path), "bias": bias}, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |