Instructions to use danielfein/raid-ce-gemma4-e4b with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use danielfein/raid-ce-gemma4-e4b with Transformers:
# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("danielfein/raid-ce-gemma4-e4b", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| from __future__ import annotations | |
| import gc | |
| import hashlib | |
| import random | |
| from collections import Counter | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| import torch | |
| import torch.nn.functional as F | |
| from tqdm.auto import tqdm | |
| from .config import TrainingConfig | |
| from .data import SourcePair | |
| from .modeling import ( | |
| ModelBundle, | |
| _model_forward, | |
| build_prompt, | |
| cosine_with_floor, | |
| encode_response, | |
| ) | |
| class TrainedTokenArtifacts: | |
| embedding: torch.Tensor | |
| loss_history: list[float] | |
| best_val_accuracy: float | |
| final_val_accuracy: float | |
| secondary_embeddings: list[torch.Tensor] = None # per-layer embeddings (e.g. Gemma 4) | |
| def _ref_cache_key(model_name: str, prompt: str, pair_ids: list[str], key: str) -> str: | |
| h = hashlib.sha256() | |
| h.update(model_name.encode()) | |
| h.update(prompt.encode()) | |
| h.update(key.encode()) | |
| for pair_id in pair_ids: | |
| h.update(pair_id.encode()) | |
| return h.hexdigest()[:16] | |
| def _source_balance_key(pair: SourcePair) -> str: | |
| source_id = str(pair.source_id or "").strip() | |
| if source_id: | |
| return source_id | |
| human_text = str(pair.human_text or "").strip() | |
| if human_text: | |
| digest = hashlib.sha256(human_text.encode("utf-8")).hexdigest()[:16] | |
| return f"human_text::{digest}" | |
| return f"pair::{pair.pair_id}" | |
| def _source_balance_weights(train_pairs: list[SourcePair]) -> torch.Tensor: | |
| keys = [_source_balance_key(pair) for pair in train_pairs] | |
| counts = Counter(keys) | |
| if not counts: | |
| return torch.ones(len(train_pairs), dtype=torch.float32) | |
| scale = len(train_pairs) / len(counts) | |
| weights = torch.tensor([scale / counts[key] for key in keys], dtype=torch.float32) | |
| print( | |
| "[source-balance] " | |
| f"pairs={len(train_pairs)} sources={len(counts)} " | |
| f"mean_weight={weights.mean().item():.4f} " | |
| f"min_weight={weights.min().item():.4f} " | |
| f"max_weight={weights.max().item():.4f}", | |
| flush=True, | |
| ) | |
| return weights | |
| def _pack_encoded_batch( | |
| bundle: ModelBundle, | |
| encoded_list: list[tuple[torch.Tensor, int]], | |
| ) -> tuple[torch.Tensor, torch.Tensor, list[int]]: | |
| pad_id = bundle.tokenizer.pad_token_id | |
| device = bundle.model.device | |
| max_len = max(ids.shape[-1] for ids, _ in encoded_list) | |
| input_ids = torch.full((len(encoded_list), max_len), pad_id, device=device) | |
| attention_mask = torch.zeros((len(encoded_list), max_len), dtype=torch.long, device=device) | |
| prompt_lens: list[int] = [] | |
| for batch_index, (ids, prompt_len) in enumerate(encoded_list): | |
| seq_len = ids.shape[-1] | |
| offset = max_len - seq_len | |
| input_ids[batch_index, offset:] = ids.squeeze().to(device) | |
| attention_mask[batch_index, offset:] = 1 | |
| prompt_lens.append(prompt_len + offset) | |
| return input_ids, attention_mask, prompt_lens | |
| def batch_sequence_logprobs( | |
| bundle: ModelBundle, | |
| encoded_list: list[tuple[torch.Tensor, int]], | |
| ) -> torch.Tensor: | |
| input_ids, attention_mask, prompt_lens = _pack_encoded_batch(bundle, encoded_list) | |
| logits = _model_forward(bundle, input_ids, attention_mask).logits | |
| results: list[torch.Tensor] = [] | |
| for batch_index, prompt_len in enumerate(prompt_lens): | |
| row_logits = logits[batch_index, prompt_len - 1 : -1] | |
| targets = input_ids[batch_index, prompt_len:] | |
| # Accumulate response likelihoods in fp32. Summing in bf16 visibly | |
| # quantizes long-sequence scores and destroys fine-grained rankings. | |
| row_log_probs = F.log_softmax(row_logits.float(), dim=-1) | |
| token_logps = row_log_probs.gather(-1, targets.unsqueeze(-1)).squeeze(-1) | |
| results.append(token_logps.sum()) | |
| return torch.stack(results) | |
| def batched_sequence_logprobs( | |
| bundle: ModelBundle, | |
| encoded_list: list[tuple[torch.Tensor, int]], | |
| *, | |
| batch_size: int, | |
| desc: str, | |
| checkpoint_path: Path | None = None, | |
| checkpoint_every: int = 1000, | |
| ) -> torch.Tensor: | |
| n = len(encoded_list) | |
| total_batches = (n + batch_size - 1) // batch_size | |
| results = torch.zeros(n) | |
| start_batch = 0 | |
| if checkpoint_path is not None and checkpoint_path.exists(): | |
| ckpt = torch.load(checkpoint_path, map_location="cpu", weights_only=True) | |
| if ckpt["n"] == n: | |
| results = ckpt["results"] | |
| start_batch = ckpt["next_batch"] | |
| if start_batch >= total_batches: | |
| return results | |
| for batch_idx in tqdm(range(start_batch, total_batches), desc=desc, initial=start_batch, total=total_batches): | |
| start = batch_idx * batch_size | |
| batch = encoded_list[start : start + batch_size] | |
| results[start : start + len(batch)] = batch_sequence_logprobs(bundle, batch).detach().cpu() | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| if checkpoint_path is not None and (batch_idx + 1) % checkpoint_every == 0: | |
| torch.save({"n": n, "results": results, "next_batch": batch_idx + 1}, checkpoint_path) | |
| return results | |
| def _load_or_compute_ref( | |
| bundle: ModelBundle, | |
| encoded_list: list[tuple[torch.Tensor, int]], | |
| *, | |
| batch_size: int, | |
| desc: str, | |
| cache_dir: Path, | |
| cache_key: str, | |
| ) -> torch.Tensor: | |
| cache_dir.mkdir(parents=True, exist_ok=True) | |
| cache_path = cache_dir / f"{cache_key}.pt" | |
| ckpt_path = cache_dir / f"{cache_key}.ckpt.pt" | |
| if cache_path.exists(): | |
| cached = torch.load(cache_path, map_location="cpu", weights_only=True) | |
| if cached.shape[0] == len(encoded_list): | |
| return cached | |
| with torch.no_grad(): | |
| result = batched_sequence_logprobs( | |
| bundle, encoded_list, batch_size=batch_size, desc=desc, | |
| checkpoint_path=ckpt_path, checkpoint_every=1000, | |
| ) | |
| torch.save(result, cache_path) | |
| if ckpt_path.exists(): | |
| ckpt_path.unlink() | |
| return result | |
| def encode_all( | |
| bundle: ModelBundle, | |
| pairs: list[SourcePair], | |
| *, | |
| prompt: str, | |
| chosen_key: str, | |
| rejected_key: str, | |
| ) -> tuple[list[tuple[torch.Tensor, int]], list[tuple[torch.Tensor, int]]]: | |
| chosen = [] | |
| rejected = [] | |
| for pair in tqdm(pairs, desc="encode chosen"): | |
| ids, prompt_len = encode_response(bundle, prompt, getattr(pair, chosen_key)) | |
| chosen.append((ids.cpu(), prompt_len)) | |
| for pair in tqdm(pairs, desc="encode rejected"): | |
| ids, prompt_len = encode_response(bundle, prompt, getattr(pair, rejected_key)) | |
| rejected.append((ids.cpu(), prompt_len)) | |
| return chosen, rejected | |
| def eval_pair_accuracy( | |
| bundle: ModelBundle, | |
| pairs: list[SourcePair], | |
| *, | |
| prompt: str, | |
| chosen_key: str, | |
| rejected_key: str, | |
| n: int, | |
| batch_size: int, | |
| ) -> float: | |
| subset = pairs[: min(n, len(pairs))] | |
| if not subset: | |
| return 0.0 | |
| chosen_enc = [] | |
| rejected_enc = [] | |
| for pair in subset: | |
| chosen_ids, chosen_prompt_len = encode_response(bundle, prompt, getattr(pair, chosen_key)) | |
| rejected_ids, rejected_prompt_len = encode_response(bundle, prompt, getattr(pair, rejected_key)) | |
| chosen_enc.append((chosen_ids.cpu(), chosen_prompt_len)) | |
| rejected_enc.append((rejected_ids.cpu(), rejected_prompt_len)) | |
| correct = 0 | |
| for start in range(0, len(subset), batch_size): | |
| chosen_logps = batch_sequence_logprobs(bundle, chosen_enc[start : start + batch_size]) | |
| rejected_logps = batch_sequence_logprobs(bundle, rejected_enc[start : start + batch_size]) | |
| correct += int((chosen_logps > rejected_logps).sum().item()) | |
| return correct / len(subset) | |
| def train_single_token( | |
| bundle: ModelBundle, | |
| train_pairs: list[SourcePair], | |
| holdout_pairs: list[SourcePair], | |
| *, | |
| token: str, | |
| learning_rate: float, | |
| chosen_key: str, | |
| rejected_key: str, | |
| config: TrainingConfig, | |
| cache_dir: Path, | |
| seed: int, | |
| precomputed_ref_chosen: torch.Tensor | None = None, | |
| precomputed_ref_rejected: torch.Tensor | None = None, | |
| ) -> tuple[TrainedTokenArtifacts, torch.Tensor, torch.Tensor]: | |
| gc.collect() | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| token_id = bundle.tokenizer.convert_tokens_to_ids(token) | |
| # Warm-start: overwrite new token rows from a neutral existing vocab token | |
| # (must happen BEFORE ref cache so ref logprobs reflect the warm start) | |
| if config.neutral_word is not None: | |
| nid = bundle.tokenizer.convert_tokens_to_ids(config.neutral_word) | |
| unk = bundle.tokenizer.unk_token_id | |
| if nid is None or nid == unk: | |
| nid = bundle.tokenizer.convert_tokens_to_ids("text") or 0 | |
| print(f"[{token}] warm-starting from '{config.neutral_word}' (id={nid})") | |
| input_emb_ws = bundle.model.get_input_embeddings() | |
| with torch.no_grad(): | |
| input_emb_ws.weight[token_id] = input_emb_ws.weight[nid].clone() | |
| for _, m in bundle.model.named_modules(): | |
| if (isinstance(m, torch.nn.Embedding) | |
| and m is not input_emb_ws | |
| and m.weight.shape[0] == len(bundle.tokenizer)): | |
| m.weight[token_id] = m.weight[nid].clone() | |
| for parameter in bundle.model.parameters(): | |
| parameter.requires_grad = False | |
| input_emb = bundle.model.get_input_embeddings() | |
| input_emb.weight.requires_grad = True | |
| # Gemma 4's auxiliary per-layer token table must be resized and mean-filled, | |
| # but its new rows remain frozen. Training it diverts signal away from the | |
| # tied input/output row and degrades the verbalization channel. | |
| secondary_embs: list[torch.nn.Embedding] = [] | |
| for name, module in bundle.model.named_modules(): | |
| if (isinstance(module, torch.nn.Embedding) | |
| and module is not input_emb | |
| and module.weight.shape[0] == len(bundle.tokenizer)): | |
| expected_mean = module.weight.data[:bundle.initial_tokenizer_len].mean( | |
| dim=0, dtype=torch.float32 | |
| ).to(dtype=module.weight.dtype) | |
| if not torch.equal(module.weight.data[token_id], expected_mean): | |
| raise RuntimeError( | |
| f"Secondary embedding {name} token row is not mean-initialized" | |
| ) | |
| module.weight.requires_grad = False | |
| secondary_embs.append(module) | |
| if secondary_embs: | |
| print( | |
| f"[{token}] freezing {len(secondary_embs)} secondary " | |
| "embedding(s) at mean initialization" | |
| ) | |
| optimizer = torch.optim.AdamW( | |
| [{"params": [input_emb.weight], "lr": learning_rate, | |
| "_lr_scale": 1.0}], | |
| weight_decay=0.0, | |
| ) | |
| prompt = build_prompt(bundle, token) | |
| chosen_enc, rejected_enc = encode_all( | |
| bundle, | |
| train_pairs, | |
| prompt=prompt, | |
| chosen_key=chosen_key, | |
| rejected_key=rejected_key, | |
| ) | |
| train_ids = [pair.pair_id for pair in train_pairs] | |
| pair_weights = _source_balance_weights(train_pairs) if config.source_balance_by_source else None | |
| # Cache key suffix includes neutral_word so warm-start and mean-init caches | |
| # never collide (neutral_word=None → empty suffix = backward-compatible) | |
| ws_suffix = f"_{config.neutral_word}" if config.neutral_word else "" | |
| if precomputed_ref_chosen is not None: | |
| print(f"[{token}] reusing precomputed ref_chosen") | |
| ref_chosen = precomputed_ref_chosen | |
| else: | |
| ref_chosen = _load_or_compute_ref( | |
| bundle, | |
| chosen_enc, | |
| batch_size=config.ref_cache_batch_size, | |
| desc=f"[{token}] ref chosen", | |
| cache_dir=cache_dir, | |
| cache_key=_ref_cache_key(bundle.config.model.model_name, prompt, train_ids, f"train_{chosen_key}{ws_suffix}"), | |
| ) | |
| if precomputed_ref_rejected is not None: | |
| print(f"[{token}] reusing precomputed ref_rejected") | |
| ref_rejected = precomputed_ref_rejected | |
| else: | |
| ref_rejected = _load_or_compute_ref( | |
| bundle, | |
| rejected_enc, | |
| batch_size=config.ref_cache_batch_size, | |
| desc=f"[{token}] ref rejected", | |
| cache_dir=cache_dir, | |
| cache_key=_ref_cache_key(bundle.config.model.model_name, prompt, train_ids, f"train_{rejected_key}{ws_suffix}"), | |
| ) | |
| total_steps = config.train_steps if config.train_steps is not None else len(train_pairs) // config.batch_size | |
| if total_steps == 0: | |
| raise ValueError(f"Not enough train pairs ({len(train_pairs)}) for batch_size={config.batch_size}.") | |
| rng = random.Random(seed) | |
| order = list(range(len(train_pairs))) | |
| rng.shuffle(order) | |
| total_needed = total_steps * config.batch_size | |
| if total_needed > len(order): | |
| order = order * (total_needed // len(order) + 1) | |
| order = order[:total_needed] | |
| best_val_accuracy = 0.0 | |
| final_val_accuracy = eval_pair_accuracy( | |
| bundle, | |
| holdout_pairs, | |
| prompt=prompt, | |
| chosen_key=chosen_key, | |
| rejected_key=rejected_key, | |
| n=config.eval_subset_size, | |
| batch_size=config.ref_cache_batch_size, | |
| ) | |
| best_val_accuracy = max(best_val_accuracy, final_val_accuracy) | |
| loss_history: list[float] = [] | |
| progress = tqdm(range(1, total_steps + 1), desc=f"[{token}] train") | |
| for step in progress: | |
| optimizer.zero_grad(set_to_none=True) | |
| batch_ids = order[(step - 1) * config.batch_size : step * config.batch_size] | |
| batch_chosen = [chosen_enc[pair_index] for pair_index in batch_ids] | |
| batch_rejected = [rejected_enc[pair_index] for pair_index in batch_ids] | |
| bundle.model.train() | |
| chosen_logps = batch_sequence_logprobs(bundle, batch_chosen) | |
| rejected_logps = batch_sequence_logprobs(bundle, batch_rejected) | |
| bundle.model.eval() | |
| ref_c = ref_chosen[batch_ids].to(bundle.model.device) | |
| ref_r = ref_rejected[batch_ids].to(bundle.model.device) | |
| dpo_losses = -F.logsigmoid(config.beta * ((chosen_logps - rejected_logps) + (ref_c - ref_r))) | |
| apo_losses = -F.logsigmoid(config.beta * (chosen_logps - ref_c)) | |
| per_example_losses = dpo_losses + config.apo_alpha * apo_losses | |
| if pair_weights is not None: | |
| weights = pair_weights[batch_ids].to(bundle.model.device) | |
| loss = (per_example_losses * weights).mean() | |
| else: | |
| loss = per_example_losses.mean() | |
| loss.backward() | |
| step_loss = float(loss.item()) | |
| # Mask grads to token_id row only + clip just that row (avoids clip_grad_norm_ on huge tensors) | |
| with torch.no_grad(): | |
| for emb in [input_emb] + secondary_embs: | |
| if emb.weight.grad is not None: | |
| row = emb.weight.grad[token_id] | |
| row_norm = row.norm() | |
| if row_norm > 1.0: | |
| row.mul_(1.0 / row_norm) | |
| emb.weight.grad[:token_id].zero_() | |
| emb.weight.grad[token_id + 1:].zero_() | |
| lr = cosine_with_floor( | |
| step - 1, | |
| total_steps, | |
| learning_rate, | |
| min_lr=config.min_learning_rate, | |
| warmup_steps=config.warmup_steps, | |
| ) | |
| # For secondary embeddings, scale LR proportionally | |
| for group in optimizer.param_groups: | |
| scale = group.get("_lr_scale", 1.0) | |
| group["lr"] = lr * scale | |
| optimizer.step() | |
| loss_history.append(step_loss) | |
| progress.set_postfix(loss=f"{step_loss:.4f}", lr=f"{lr:.2e}") | |
| if step % config.eval_every_steps == 0 or step == total_steps: | |
| final_val_accuracy = eval_pair_accuracy( | |
| bundle, | |
| holdout_pairs, | |
| prompt=prompt, | |
| chosen_key=chosen_key, | |
| rejected_key=rejected_key, | |
| n=config.eval_subset_size, | |
| batch_size=config.ref_cache_batch_size, | |
| ) | |
| best_val_accuracy = max(best_val_accuracy, final_val_accuracy) | |
| msg = f"[{token}] step {step:>5d} val_acc: {final_val_accuracy:.3f}" | |
| progress.write(msg) | |
| print(msg, flush=True) | |
| embedding = input_emb.weight[token_id].detach().cpu() | |
| secondary_embeddings = [e.weight[token_id].detach().cpu() for e in secondary_embs] if secondary_embs else None | |
| gc.collect() | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| return TrainedTokenArtifacts( | |
| embedding=embedding, | |
| loss_history=loss_history, | |
| best_val_accuracy=best_val_accuracy, | |
| final_val_accuracy=final_val_accuracy, | |
| secondary_embeddings=secondary_embeddings, | |
| ), ref_chosen, ref_rejected | |