GitLab Repo

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