amachine.am_transformers.am_eval
1import math 2import contextlib 3from typing import Any 4import torch 5from .am_basis_eval import PrivilegedBasisEvaluator 6from transformers import PreTrainedTokenizerFast 7 8def _strip_trailing_eos( 9 input_ids: list[int], 10 eos_id: int | None) -> list[int]: 11 12 if eos_id is None: 13 return input_ids 14 end = len(input_ids) 15 while end > 0 and input_ids[end - 1] == eos_id: 16 end -= 1 17 return input_ids[:end] 18 19@torch.inference_mode() 20def generate_test_prompts( 21 model: Any, 22 raw_tokenizer: Any, 23 pad_token : str | None, 24 eos_token : str | None, 25 unk_token : str | None, 26 test_prompts: list[str] | None, 27 gen_len: int | None, 28 device: torch.device, 29 ptdtype: torch.dtype, 30 use_amp: bool, 31) -> list[dict[str, str]]: 32 33 hf_tok = PreTrainedTokenizerFast( 34 tokenizer_object=raw_tokenizer, 35 pad_token=pad_token, 36 eos_token=eos_token, 37 unk_token=unk_token 38 ) 39 40 if not test_prompts: 41 return [] 42 43 if not hasattr(model, "generate"): 44 return [] 45 46 FALLBACK_GEN_LEN = 128 47 if gen_len is None or gen_len <= 0: 48 gen_len = FALLBACK_GEN_LEN 49 50 CTX_LIMIT = 4096 51 max_context_length = ( 52 getattr(model.config, "max_position_embeddings", None) 53 or getattr(model.config, "n_positions", None) 54 or getattr(model.config, "n_ctx", None) 55 or CTX_LIMIT 56 ) 57 58 if hf_tok.pad_token_id is None: 59 hf_tok.pad_token = hf_tok.eos_token 60 61 prior_padding_side = hf_tok.padding_side 62 prior_truncation_side = hf_tok.truncation_side 63 64 hf_tok.padding_side = "left" 65 hf_tok.truncation_side = "right" 66 67 try: 68 69 eos_id = hf_tok.eos_token_id 70 assert isinstance( eos_id, int ) 71 72 per_prompt_ids = [ 73 _strip_trailing_eos( hf_tok( p, add_special_tokens=True)["input_ids"], eos_id ) 74 for p in test_prompts 75 ] 76 77 enc = hf_tok.pad( 78 {"input_ids": per_prompt_ids}, 79 return_tensors="pt", 80 padding=True, 81 ).to(device) 82 83 prompt_len = enc["input_ids"].shape[1] 84 85 if max_context_length is not None: 86 room = max_context_length - prompt_len 87 if room <= 0: 88 # prompts alone already fill the context window 89 return [ 90 {"prompt": p, "continuation": ""} 91 for p in test_prompts 92 ] 93 gen_len = min( gen_len, room ) 94 95 with torch.autocast(device_type=device.type, dtype=ptdtype, enabled=use_amp): 96 gen_ids = model.generate( 97 **enc, 98 max_new_tokens=gen_len, 99 do_sample=False, 100 num_beams=1, 101 pad_token_id=hf_tok.pad_token_id, 102 eos_token_id=hf_tok.eos_token_id, 103 use_cache=False 104 ) 105 106 continuations = hf_tok.batch_decode( 107 gen_ids[:, prompt_len:], 108 skip_special_tokens=True 109 ) 110 111 return [ 112 {"prompt": p, "continuation": c} 113 for p, c in zip(test_prompts, continuations) 114 ] 115 116 except (RuntimeError, torch.cuda.OutOfMemoryError) as e: 117 return [{"prompt": "generation_error", "continuation": str(e)}] 118 119 finally: 120 hf_tok.padding_side = prior_padding_side 121 hf_tok.truncation_side = prior_truncation_side 122 123def run_eval( 124 eval_loader: Any, 125 model: Any, 126 training_config: Any, 127 device: torch.device, 128 ptdtype: torch.dtype, 129 use_amp: bool, 130 pad_id: int, 131 pad_token: str | None, 132 eos_token: str | None, 133 unk_token: str | None, 134 tokenizer: Any | None = None, 135) -> dict[str, Any] | None: 136 137 if eval_loader is None: 138 return None 139 140 model.eval() 141 142 try: 143 144 eval_basis = getattr( training_config, "evaluate_basis_independence", False ) 145 basis_evaluator = PrivilegedBasisEvaluator( 146 model, 147 pad_id=pad_id, 148 compute_effective_rank=training_config.compute_effective_rank, 149 compute_channel_stats=training_config.compute_channel_stats 150 ) if eval_basis else None 151 ctx = basis_evaluator if basis_evaluator is not None else contextlib.nullcontext() 152 153 with torch.inference_mode(), ctx: 154 155 total_loss = total_tokens = n_batches = 0 156 eval_iter = iter(eval_loader) 157 158 extra_term_names = getattr(training_config, "extra_loss_terms", None) or [] 159 extra_term_sums = {name: 0.0 for name in extra_term_names} 160 161 for _ in range(training_config.eval_steps): 162 try: 163 batch = next(eval_iter) 164 except StopIteration: 165 break 166 167 input_ids = batch["input_ids"].to(device, non_blocking=True) 168 labels = input_ids.masked_fill(input_ids == pad_id, -100) 169 170 # Tell the hooks which tokens are padding *before* the forward runs. 171 if basis_evaluator is not None: 172 basis_evaluator.observe( input_ids ) 173 174 with torch.autocast(device_type=device.type, dtype=ptdtype, enabled=use_amp): 175 outputs = model(input_ids=input_ids, labels=labels) 176 177 batch_tokens = (labels[..., 1:] != -100).sum().item() 178 total_loss += outputs.loss.item() * batch_tokens 179 total_tokens += batch_tokens 180 n_batches += 1 181 182 for name in extra_term_names: 183 term_val = getattr(outputs, name, None) 184 if term_val is not None: 185 if isinstance( term_val, float ) : 186 extra_term_sums[name] += term_val 187 else : 188 extra_term_sums[name] += term_val.detach().item() 189 190 if total_tokens == 0: 191 model.train() 192 return None 193 194 eval_loss = total_loss / total_tokens 195 196 results = { 197 "eval_loss": eval_loss, 198 "eval_perplexity": math.exp(min(eval_loss, 20)), 199 "eval_batches": n_batches, 200 } 201 202 eval_extra_terms = {name: s / max(n_batches, 1) for name, s in extra_term_sums.items()} 203 if eval_extra_terms: 204 results["extra_loss_terms"] = eval_extra_terms 205 206 if basis_evaluator is not None: 207 res = basis_evaluator.finalize() 208 results.update( res ) 209 210 test_prompts = getattr(training_config, "test_prompts", None) 211 if test_prompts and tokenizer is not None: 212 results["test_generations"] = generate_test_prompts( 213 model=model, 214 raw_tokenizer=tokenizer, 215 pad_token=pad_token, 216 eos_token=eos_token, 217 unk_token=unk_token, 218 test_prompts=test_prompts, 219 gen_len=training_config.test_prompt_gen_len, 220 device=device, 221 ptdtype=ptdtype, 222 use_amp=use_amp, 223 ) 224 225 226 finally : 227 228 model.train() 229 230 return results
@torch.inference_mode()
def
generate_test_prompts( model: Any, raw_tokenizer: Any, pad_token: str | None, eos_token: str | None, unk_token: str | None, test_prompts: list[str] | None, gen_len: int | None, device: torch.device, ptdtype: torch.dtype, use_amp: bool) -> list[dict[str, str]]:
20@torch.inference_mode() 21def generate_test_prompts( 22 model: Any, 23 raw_tokenizer: Any, 24 pad_token : str | None, 25 eos_token : str | None, 26 unk_token : str | None, 27 test_prompts: list[str] | None, 28 gen_len: int | None, 29 device: torch.device, 30 ptdtype: torch.dtype, 31 use_amp: bool, 32) -> list[dict[str, str]]: 33 34 hf_tok = PreTrainedTokenizerFast( 35 tokenizer_object=raw_tokenizer, 36 pad_token=pad_token, 37 eos_token=eos_token, 38 unk_token=unk_token 39 ) 40 41 if not test_prompts: 42 return [] 43 44 if not hasattr(model, "generate"): 45 return [] 46 47 FALLBACK_GEN_LEN = 128 48 if gen_len is None or gen_len <= 0: 49 gen_len = FALLBACK_GEN_LEN 50 51 CTX_LIMIT = 4096 52 max_context_length = ( 53 getattr(model.config, "max_position_embeddings", None) 54 or getattr(model.config, "n_positions", None) 55 or getattr(model.config, "n_ctx", None) 56 or CTX_LIMIT 57 ) 58 59 if hf_tok.pad_token_id is None: 60 hf_tok.pad_token = hf_tok.eos_token 61 62 prior_padding_side = hf_tok.padding_side 63 prior_truncation_side = hf_tok.truncation_side 64 65 hf_tok.padding_side = "left" 66 hf_tok.truncation_side = "right" 67 68 try: 69 70 eos_id = hf_tok.eos_token_id 71 assert isinstance( eos_id, int ) 72 73 per_prompt_ids = [ 74 _strip_trailing_eos( hf_tok( p, add_special_tokens=True)["input_ids"], eos_id ) 75 for p in test_prompts 76 ] 77 78 enc = hf_tok.pad( 79 {"input_ids": per_prompt_ids}, 80 return_tensors="pt", 81 padding=True, 82 ).to(device) 83 84 prompt_len = enc["input_ids"].shape[1] 85 86 if max_context_length is not None: 87 room = max_context_length - prompt_len 88 if room <= 0: 89 # prompts alone already fill the context window 90 return [ 91 {"prompt": p, "continuation": ""} 92 for p in test_prompts 93 ] 94 gen_len = min( gen_len, room ) 95 96 with torch.autocast(device_type=device.type, dtype=ptdtype, enabled=use_amp): 97 gen_ids = model.generate( 98 **enc, 99 max_new_tokens=gen_len, 100 do_sample=False, 101 num_beams=1, 102 pad_token_id=hf_tok.pad_token_id, 103 eos_token_id=hf_tok.eos_token_id, 104 use_cache=False 105 ) 106 107 continuations = hf_tok.batch_decode( 108 gen_ids[:, prompt_len:], 109 skip_special_tokens=True 110 ) 111 112 return [ 113 {"prompt": p, "continuation": c} 114 for p, c in zip(test_prompts, continuations) 115 ] 116 117 except (RuntimeError, torch.cuda.OutOfMemoryError) as e: 118 return [{"prompt": "generation_error", "continuation": str(e)}] 119 120 finally: 121 hf_tok.padding_side = prior_padding_side 122 hf_tok.truncation_side = prior_truncation_side
def
run_eval( eval_loader: Any, model: Any, training_config: Any, device: torch.device, ptdtype: torch.dtype, use_amp: bool, pad_id: int, pad_token: str | None, eos_token: str | None, unk_token: str | None, tokenizer: typing.Any | None = None) -> dict[str, typing.Any] | None:
124def run_eval( 125 eval_loader: Any, 126 model: Any, 127 training_config: Any, 128 device: torch.device, 129 ptdtype: torch.dtype, 130 use_amp: bool, 131 pad_id: int, 132 pad_token: str | None, 133 eos_token: str | None, 134 unk_token: str | None, 135 tokenizer: Any | None = None, 136) -> dict[str, Any] | None: 137 138 if eval_loader is None: 139 return None 140 141 model.eval() 142 143 try: 144 145 eval_basis = getattr( training_config, "evaluate_basis_independence", False ) 146 basis_evaluator = PrivilegedBasisEvaluator( 147 model, 148 pad_id=pad_id, 149 compute_effective_rank=training_config.compute_effective_rank, 150 compute_channel_stats=training_config.compute_channel_stats 151 ) if eval_basis else None 152 ctx = basis_evaluator if basis_evaluator is not None else contextlib.nullcontext() 153 154 with torch.inference_mode(), ctx: 155 156 total_loss = total_tokens = n_batches = 0 157 eval_iter = iter(eval_loader) 158 159 extra_term_names = getattr(training_config, "extra_loss_terms", None) or [] 160 extra_term_sums = {name: 0.0 for name in extra_term_names} 161 162 for _ in range(training_config.eval_steps): 163 try: 164 batch = next(eval_iter) 165 except StopIteration: 166 break 167 168 input_ids = batch["input_ids"].to(device, non_blocking=True) 169 labels = input_ids.masked_fill(input_ids == pad_id, -100) 170 171 # Tell the hooks which tokens are padding *before* the forward runs. 172 if basis_evaluator is not None: 173 basis_evaluator.observe( input_ids ) 174 175 with torch.autocast(device_type=device.type, dtype=ptdtype, enabled=use_amp): 176 outputs = model(input_ids=input_ids, labels=labels) 177 178 batch_tokens = (labels[..., 1:] != -100).sum().item() 179 total_loss += outputs.loss.item() * batch_tokens 180 total_tokens += batch_tokens 181 n_batches += 1 182 183 for name in extra_term_names: 184 term_val = getattr(outputs, name, None) 185 if term_val is not None: 186 if isinstance( term_val, float ) : 187 extra_term_sums[name] += term_val 188 else : 189 extra_term_sums[name] += term_val.detach().item() 190 191 if total_tokens == 0: 192 model.train() 193 return None 194 195 eval_loss = total_loss / total_tokens 196 197 results = { 198 "eval_loss": eval_loss, 199 "eval_perplexity": math.exp(min(eval_loss, 20)), 200 "eval_batches": n_batches, 201 } 202 203 eval_extra_terms = {name: s / max(n_batches, 1) for name, s in extra_term_sums.items()} 204 if eval_extra_terms: 205 results["extra_loss_terms"] = eval_extra_terms 206 207 if basis_evaluator is not None: 208 res = basis_evaluator.finalize() 209 results.update( res ) 210 211 test_prompts = getattr(training_config, "test_prompts", None) 212 if test_prompts and tokenizer is not None: 213 results["test_generations"] = generate_test_prompts( 214 model=model, 215 raw_tokenizer=tokenizer, 216 pad_token=pad_token, 217 eos_token=eos_token, 218 unk_token=unk_token, 219 test_prompts=test_prompts, 220 gen_len=training_config.test_prompt_gen_len, 221 device=device, 222 ptdtype=ptdtype, 223 use_amp=use_amp, 224 ) 225 226 227 finally : 228 229 model.train() 230 231 return results