Zero-Shot Classification
Safetensors
PEFT
English
openjev
classification
decision-model
listwise
gemma4
research
Instructions to use bambamdevs/openjev-e4b with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use bambamdevs/openjev-e4b with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
Download openjev/formatting.py from bambamdevs/openjev-e4b: direct link, hf CLI and curl.
- Browser
- Download file 8.83 kB
-
https://huggingface.co/bambamdevs/openjev-e4b/resolve/main/openjev/formatting.py
- Command line
-
hf download hf://bambamdevs/openjev-e4b/openjev/formatting.py
-
curl -L -o formatting.py https://huggingface.co/bambamdevs/openjev-e4b/resolve/main/openjev/formatting.py
8.83 kB
| from __future__ import annotations | |
| import json | |
| import torch | |
| def encode_piece(tok, text: str) -> list[int]: | |
| return tok(text, add_special_tokens=False).input_ids | |
| def state_to_text(state) -> str: | |
| if isinstance(state, str): | |
| return state | |
| return json.dumps(state, ensure_ascii=False, sort_keys=True, separators=(",", ":")) | |
| def _clip_head_tail(ids: list[int], limit: int) -> list[int]: | |
| """Deterministically keep evidence from both ends of a long field.""" | |
| if limit <= 0: | |
| return [] | |
| if len(ids) <= limit: | |
| return ids | |
| head = (limit + 1) // 2 | |
| tail = limit - head | |
| return ids[:head] + (ids[-tail:] if tail else []) | |
| def _assemble(tok, state_ids, q_ids, option_ids, max_length: int): | |
| ids = [tok.bos_token_id] if tok.bos_token_id is not None else [] | |
| ids += encode_piece(tok, "State:\n") | |
| ids += state_ids | |
| ids += encode_piece(tok, "\n\nQuestion:\n") | |
| ids += q_ids | |
| ids += encode_piece(tok, "\n\nOptions:\n") | |
| option_positions = [] | |
| option_spans = [] | |
| for idx, opt_ids in enumerate(option_ids): | |
| ids += encode_piece(tok, f"- [{idx}] ") | |
| span_start = len(ids) | |
| ids += opt_ids | |
| span_end = len(ids) | |
| option_spans.append((span_start, span_end)) | |
| # Represent each option by its final semantic token, after it has seen | |
| # the full option text but before the newline delimiter. | |
| option_positions.append(len(ids) - 1) | |
| ids += encode_piece(tok, "\n") | |
| ids += encode_piece(tok, "\nDecision:") | |
| decide_position = len(ids) - 1 | |
| if len(ids) > max_length: | |
| return None | |
| return ids, option_positions, decide_position, option_spans | |
| def pack_question(tok, state, question: dict, max_length: int): | |
| """Pack a decision example without dropping long states. | |
| Priority order is structural markers/options/question first, then state. | |
| Long state is head+tail truncated. If question/options themselves are huge, | |
| they are bounded as a second-stage fallback. Returning None is reserved for | |
| structurally impossible cases (e.g. too many options for max_length). | |
| """ | |
| if max_length < 32: | |
| return None | |
| state_ids = encode_piece(tok, state_to_text(state)) | |
| q_ids = encode_piece(tok, str(question["instruction"])) | |
| option_ids = [encode_piece(tok, str(opt["text"])) for opt in question["options"]] | |
| if not option_ids or any(not x for x in option_ids): | |
| return None | |
| was_truncated = False | |
| # First discover how much room remains for the state while preserving the | |
| # complete question and option text. | |
| probe = _assemble(tok, [], q_ids, option_ids, max_length=10**9) | |
| if probe is None: | |
| return None | |
| structural_len = len(probe[0]) | |
| if structural_len <= max_length: | |
| state_budget = max_length - structural_len | |
| was_truncated = len(state_ids) > state_budget | |
| packed = _assemble(tok, _clip_head_tail(state_ids, state_budget), q_ids, option_ids, max_length) | |
| else: | |
| was_truncated = True | |
| # Extremely verbose question/tool schemas: cap semantic fields rather | |
| # than dropping the sample. Typical training examples never hit this. | |
| nopt = len(option_ids) | |
| q_cap = min(len(q_ids), max(16, max_length // 8)) | |
| # Start modestly; then shrink until the structural representation fits. | |
| opt_cap = max(8, min(96, max_length // max(8, nopt * 2))) | |
| q_fit = _clip_head_tail(q_ids, q_cap) | |
| opts_fit = [_clip_head_tail(x, opt_cap) for x in option_ids] | |
| packed = _assemble(tok, [], q_fit, opts_fit, max_length) | |
| while packed is None and (q_cap > 8 or opt_cap > 4): | |
| q_cap = max(8, q_cap // 2) | |
| opt_cap = max(4, opt_cap // 2) | |
| q_fit = _clip_head_tail(q_ids, q_cap) | |
| opts_fit = [_clip_head_tail(x, opt_cap) for x in option_ids] | |
| packed = _assemble(tok, [], q_fit, opts_fit, max_length) | |
| if packed is not None: | |
| # If the shrunken question/options leave room, fill it with state. | |
| base_len = len(packed[0]) | |
| state_budget = max(0, max_length - base_len) | |
| was_truncated = was_truncated or len(state_ids) > state_budget | |
| packed = _assemble(tok, _clip_head_tail(state_ids, state_budget), q_fit, opts_fit, max_length) | |
| if packed is None: | |
| return None | |
| ids, option_positions, decide_position, option_spans = packed | |
| return { | |
| "input_ids": ids, | |
| "option_positions": option_positions, | |
| "option_spans": option_spans, | |
| "decide_position": decide_position, | |
| "target": question.get("target_distribution"), | |
| "was_truncated": was_truncated, | |
| } | |
| def collate_packed(tok, packed: list[dict]) -> dict[str, torch.Tensor]: | |
| if not packed: | |
| raise RuntimeError("No packable examples") | |
| bsz = len(packed) | |
| max_seq = max(len(x["input_ids"]) for x in packed) | |
| max_opts = max(len(x["option_positions"]) for x in packed) | |
| pad_id = tok.pad_token_id if tok.pad_token_id is not None else 0 | |
| input_ids = torch.full((bsz, max_seq), pad_id, dtype=torch.long) | |
| attention_mask = torch.zeros((bsz, max_seq), dtype=torch.long) | |
| option_positions = torch.zeros((bsz, max_opts), dtype=torch.long) | |
| option_mask = torch.zeros((bsz, max_opts), dtype=torch.bool) | |
| decide_positions = torch.zeros((bsz,), dtype=torch.long) | |
| targets = torch.zeros((bsz, max_opts), dtype=torch.float32) | |
| option_starts = torch.zeros((bsz, max_opts), dtype=torch.long) | |
| option_ends = torch.zeros((bsz, max_opts), dtype=torch.long) | |
| has_targets = all(x.get("target") is not None for x in packed) | |
| for i, x in enumerate(packed): | |
| n = len(x["input_ids"]) | |
| m = len(x["option_positions"]) | |
| input_ids[i, :n] = torch.tensor(x["input_ids"], dtype=torch.long) | |
| attention_mask[i, :n] = 1 | |
| option_positions[i, :m] = torch.tensor(x["option_positions"], dtype=torch.long) | |
| spans = x.get("option_spans") or [(int(v), int(v)+1) for v in x["option_positions"]] | |
| option_starts[i, :m] = torch.tensor([a for a, _ in spans], dtype=torch.long) | |
| option_ends[i, :m] = torch.tensor([b for _, b in spans], dtype=torch.long) | |
| option_mask[i, :m] = True | |
| decide_positions[i] = x["decide_position"] | |
| if has_targets: | |
| targets[i, :m] = torch.tensor(x["target"], dtype=torch.float32) | |
| out = { | |
| "input_ids": input_ids, | |
| "attention_mask": attention_mask, | |
| "option_positions": option_positions, | |
| "option_starts": option_starts, | |
| "option_ends": option_ends, | |
| "option_mask": option_mask, | |
| "decide_positions": decide_positions, | |
| } | |
| if has_targets: | |
| out["targets"] = targets | |
| return out | |
| def pack_shared_request(tok, state, questions: list[dict], max_length: int): | |
| """Pack one common state prefix plus causal question suffixes. | |
| Every question sees exactly the same serialized/clipped state. This is the | |
| representation needed for safe KV-prefix reuse at inference time. | |
| """ | |
| if max_length < 32 or not questions: | |
| return None | |
| bos = [tok.bos_token_id] if tok.bos_token_id is not None else [] | |
| state_marker = encode_piece(tok, "State:\n") | |
| state_ids = encode_piece(tok, state_to_text(state)) | |
| suffixes = [] | |
| max_suffix_len = 0 | |
| for question in questions: | |
| q_ids = encode_piece(tok, str(question["instruction"])) | |
| option_ids = [encode_piece(tok, str(opt["text"])) for opt in question["options"]] | |
| if not option_ids or any(not x for x in option_ids): | |
| return None | |
| ids = encode_piece(tok, "\n\nQuestion:\n") + q_ids + encode_piece(tok, "\n\nOptions:\n") | |
| option_positions = [] | |
| for idx, opt_ids in enumerate(option_ids): | |
| ids += encode_piece(tok, f"- [{idx}] ") | |
| ids += opt_ids | |
| option_positions.append(len(ids) - 1) | |
| ids += encode_piece(tok, "\n") | |
| ids += encode_piece(tok, "\nDecision:") | |
| decide_position = len(ids) - 1 | |
| suffixes.append({ | |
| "input_ids": ids, | |
| "option_positions": option_positions, | |
| "decide_position": decide_position, | |
| }) | |
| max_suffix_len = max(max_suffix_len, len(ids)) | |
| fixed_prefix_len = len(bos) + len(state_marker) | |
| state_budget = max_length - fixed_prefix_len - max_suffix_len | |
| if state_budget < 0: | |
| return None | |
| clipped = _clip_head_tail(state_ids, state_budget) | |
| was_truncated = len(clipped) < len(state_ids) | |
| prefix_ids = bos + state_marker + clipped | |
| for s in suffixes: | |
| s["was_truncated"] = was_truncated | |
| if len(prefix_ids) + len(s["input_ids"]) > max_length: | |
| return None | |
| return {"prefix_ids": prefix_ids, "suffixes": suffixes, "was_truncated": was_truncated} | |