armaniii commited on
Commit
007cea2
·
verified ·
1 Parent(s): 52c7cc4

Model card v2: empirically verified instructions (tested on transformers 5.12/peft 0.19 and 4.38/0.7.1), corrected storage-format description, verified example outputs, dependency matrix

Browse files
Files changed (1) hide show
  1. README.md +44 -27
README.md CHANGED
@@ -11,10 +11,12 @@ tags:
11
  - claim-extraction
12
  - computational-social-science
13
  - llama
 
 
14
  - wiba
15
  ---
16
 
17
- # WIBA Claim Topic Extraction (Llama-3-8B, full model)
18
 
19
  **Topic extraction** model: given an argumentative sentence or passage, it generates the **topic being argued** (a short phrase naming the person, place, thing, entity, or idea at issue), or **`No Topic`** if the text is not an argument. The topic may be explicit in the text or implicit and inferred from context.
20
 
@@ -23,25 +25,29 @@ This is **Stage 2** of the [WIBA (What Is Being Argued?)](https://arxiv.org/abs/
23
  | Stage | Task | Model | Type |
24
  |---|---|---|---|
25
  | 1. Detect | Is this text an argument? | [armaniii/llama-3-8b-argument-detection](https://huggingface.co/armaniii/llama-3-8b-argument-detection) | LoRA adapter (sequence classification, 2 labels) |
26
- | **2. Extract** | What topic is being argued? | **this repo** | Full fine-tuned causal LM |
27
  | 3. Stance | What position does it take on the topic? | [armaniii/llama-stance-classification](https://huggingface.co/armaniii/llama-stance-classification) | LoRA adapter (sequence classification, 3 labels) |
28
 
29
  - 📄 Paper: [WIBA: What Is Being Argued? A Comprehensive Approach to Argument Mining](https://arxiv.org/abs/2405.00828)
30
  - 💻 Code: [github.com/Armaniii/WIBA](https://github.com/Armaniii/WIBA)
31
  - 🌐 Platform: [wiba.dev](https://wiba.dev)
32
 
33
- ## What this repo contains (full model, not an adapter)
34
 
35
- Unlike the detect and stance stages, this repo is a **complete, self-contained fine-tuned model** (`LlamaForCausalLM`, ~16 GB of float16 safetensors in two shards). You do **not** need to download the base Llama-3 weights or merge any adapter — `from_pretrained` on this repo alone is enough.
36
 
37
  | File | Purpose |
38
  |---|---|
39
- | `model-0000*-of-00002.safetensors` + index | Full fine-tuned weights (float16) |
40
- | `config.json` | Model config — **note: ships with a bitsandbytes 4-bit (nf4) `quantization_config` baked in** (see below) |
41
  | `generation_config.json` | Default generation settings |
42
  | `tokenizer.json`, `tokenizer_config.json`, `special_tokens_map.json` | Llama-3 tokenizer |
43
 
44
- > **Quantization note:** because `config.json` includes a `quantization_config`, `from_pretrained` will automatically load the model 4-bit quantized (~6 GB VRAM) when `bitsandbytes` is installed — this matches the production WIBA deployment. To load in full fp16 instead, strip the quantization config (snippet below).
 
 
 
 
45
 
46
  ## Quickstart
47
 
@@ -55,30 +61,16 @@ from transformers import AutoTokenizer, AutoModelForCausalLM
55
 
56
  REPO = "armaniii/llama-3-8b-claim-topic-extraction"
57
 
58
- tokenizer = AutoTokenizer.from_pretrained(REPO, use_fast=False)
59
  tokenizer.pad_token_id = tokenizer.eos_token_id
60
  tokenizer.padding_side = "left"
61
 
62
- # Loads 4-bit (nf4) automatically via the config's quantization_config (~6 GB VRAM)
63
- model = AutoModelForCausalLM.from_pretrained(
64
- REPO, torch_dtype=torch.float16, device_map="auto", low_cpu_mem_usage=True
65
- )
66
  model.eval()
67
  ```
68
 
69
- To load **full fp16** (~16 GB VRAM) instead of 4-bit:
70
-
71
- ```python
72
- from transformers import AutoConfig
73
-
74
- config = AutoConfig.from_pretrained(REPO)
75
- if hasattr(config, "quantization_config"):
76
- delattr(config, "quantization_config")
77
- model = AutoModelForCausalLM.from_pretrained(
78
- REPO, config=config, torch_dtype=torch.float16, device_map="auto"
79
- )
80
- ```
81
-
82
  ### Prompt format (must match training)
83
 
84
  The model expects the Llama-3 chat header format with the WIBA topic-extraction system prompt, and the generation cut off after a few tokens (topics are short):
@@ -115,11 +107,15 @@ def extract_topic(text: str) -> str:
115
  return tokenizer.decode(out[0, enc.input_ids.shape[1]:], skip_special_tokens=True).strip()
116
 
117
  print(extract_topic("We must act on climate change because temperatures are rising."))
118
- # -> a topic phrase, e.g. "Climate change"
119
  print(extract_topic("The weather is nice today."))
120
- # -> "No Topic"
 
 
121
  ```
122
 
 
 
123
  The original implementation uses the equivalent `pipeline("text-generation", ..., max_new_tokens=8, pad_token_id=128009)` and takes the text after the final `assistant<|end_header_id|>\n\n` marker — the function above does the same thing with `generate`.
124
 
125
  ### Output
@@ -127,6 +123,27 @@ The original implementation uses the equivalent `pipeline("text-generation", ...
127
  - An argumentative input → a short topic phrase (e.g. `Climate change`, `Gun control`)
128
  - A non-argument input → the literal string `No Topic`
129
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
130
  ## How it's used in the WIBA implementation
131
 
132
  In the WIBA serving code, this model backs the `/api/extract` endpoint at [wiba.dev](https://wiba.dev). Texts that Stage 1 classified as `Argument` are passed here to name the topic; the (text, topic) pair is then passed to Stage 3 ([stance classification](https://huggingface.co/armaniii/llama-stance-classification)) to determine whether the argument is in favor of or against that topic. For batch processing the implementation streams prompts through the pipeline with `batch_size=2` and left-padding.
 
11
  - claim-extraction
12
  - computational-social-science
13
  - llama
14
+ - 4-bit
15
+ - bitsandbytes
16
  - wiba
17
  ---
18
 
19
+ # WIBA Claim Topic Extraction (Llama-3-8B, pre-quantized 4-bit)
20
 
21
  **Topic extraction** model: given an argumentative sentence or passage, it generates the **topic being argued** (a short phrase naming the person, place, thing, entity, or idea at issue), or **`No Topic`** if the text is not an argument. The topic may be explicit in the text or implicit and inferred from context.
22
 
 
25
  | Stage | Task | Model | Type |
26
  |---|---|---|---|
27
  | 1. Detect | Is this text an argument? | [armaniii/llama-3-8b-argument-detection](https://huggingface.co/armaniii/llama-3-8b-argument-detection) | LoRA adapter (sequence classification, 2 labels) |
28
+ | **2. Extract** | What topic is being argued? | **this repo** | Fine-tuned causal LM (pre-quantized 4-bit) |
29
  | 3. Stance | What position does it take on the topic? | [armaniii/llama-stance-classification](https://huggingface.co/armaniii/llama-stance-classification) | LoRA adapter (sequence classification, 3 labels) |
30
 
31
  - 📄 Paper: [WIBA: What Is Being Argued? A Comprehensive Approach to Argument Mining](https://arxiv.org/abs/2405.00828)
32
  - 💻 Code: [github.com/Armaniii/WIBA](https://github.com/Armaniii/WIBA)
33
  - 🌐 Platform: [wiba.dev](https://wiba.dev)
34
 
35
+ ## What this repo contains (full model, stored 4-bit quantized)
36
 
37
+ This repo is a **complete, self-contained fine-tuned model** — no base download, no adapter. But unlike a normal fp16 checkpoint, the weights are **stored pre-quantized with bitsandbytes NF4** (the format the WIBA platform serves in production):
38
 
39
  | File | Purpose |
40
  |---|---|
41
+ | `model-0000*-of-00002.safetensors` + index | ~6 GB total. Linear-layer weights as packed 4-bit (uint8) with `absmax`/`quant_map` quantization metadata; embeddings and `lm_head` in float16 |
42
+ | `config.json` | Model config including the `quantization_config` (bnb NF4, blocksize 64, compute dtype fp16) that tells transformers how to load the 4-bit weights |
43
  | `generation_config.json` | Default generation settings |
44
  | `tokenizer.json`, `tokenizer_config.json`, `special_tokens_map.json` | Llama-3 tokenizer |
45
 
46
+ Practical consequences:
47
+
48
+ - **`bitsandbytes` is a hard requirement** — the checkpoint cannot be loaded without it.
49
+ - Do **not** try to remove/override `quantization_config` to get fp16: the stored weights themselves are 4-bit packed, so there is no full-precision copy in this repo. To obtain higher-precision weights, load 4-bit first and call `model.dequantize()` (see below).
50
+ - VRAM needed is only **~6 GB** — the model fits on small GPUs.
51
 
52
  ## Quickstart
53
 
 
61
 
62
  REPO = "armaniii/llama-3-8b-claim-topic-extraction"
63
 
64
+ tokenizer = AutoTokenizer.from_pretrained(REPO)
65
  tokenizer.pad_token_id = tokenizer.eos_token_id
66
  tokenizer.padding_side = "left"
67
 
68
+ # quantization_config ships in config.json — transformers loads the 4-bit
69
+ # weights automatically (CUDA GPU recommended, ~6 GB VRAM)
70
+ model = AutoModelForCausalLM.from_pretrained(REPO, device_map="auto", low_cpu_mem_usage=True)
 
71
  model.eval()
72
  ```
73
 
 
 
 
 
 
 
 
 
 
 
 
 
 
74
  ### Prompt format (must match training)
75
 
76
  The model expects the Llama-3 chat header format with the WIBA topic-extraction system prompt, and the generation cut off after a few tokens (topics are short):
 
107
  return tokenizer.decode(out[0, enc.input_ids.shape[1]:], skip_special_tokens=True).strip()
108
 
109
  print(extract_topic("We must act on climate change because temperatures are rising."))
110
+ # -> climate change
111
  print(extract_topic("The weather is nice today."))
112
+ # -> No Topic
113
+ print(extract_topic("Abortion should remain legal because bodily autonomy is a fundamental right."))
114
+ # -> abortion
115
  ```
116
 
117
+ (Outputs above are actual verified predictions, not illustrations.)
118
+
119
  The original implementation uses the equivalent `pipeline("text-generation", ..., max_new_tokens=8, pad_token_id=128009)` and takes the text after the final `assistant<|end_header_id|>\n\n` marker — the function above does the same thing with `generate`.
120
 
121
  ### Output
 
123
  - An argumentative input → a short topic phrase (e.g. `Climate change`, `Gun control`)
124
  - A non-argument input → the literal string `No Topic`
125
 
126
+ ## Getting full-precision weights
127
+
128
+ The repo stores no fp16 copy, but you can dequantize after loading (needs enough memory for the fp16 model, ~16 GB):
129
+
130
+ ```python
131
+ model = AutoModelForCausalLM.from_pretrained(REPO, device_map="auto")
132
+ model = model.dequantize() # bnb 4-bit -> floating point
133
+ ```
134
+
135
+ ## Tested configurations
136
+
137
+ | Stack | Versions | Status |
138
+ |---|---|---|
139
+ | Modern (2026) | torch 2.5.1, transformers 5.12.0, accelerate 1.14.0, bitsandbytes 0.49.2 | ✅ verified (4-bit load, generation, and `dequantize()` path) |
140
+
141
+ Notes:
142
+ - Without `bitsandbytes` installed, `from_pretrained` raises immediately (the checkpoint is pre-quantized).
143
+ - Attempting to load with the `quantization_config` removed fails with shape errors (`ckpt torch.Size([8388608, 1]) vs model torch.Size([4096, 4096])`) — the stored weights really are 4-bit packed.
144
+ - CPU-only machines: the 4-bit load works (~4 GB RAM, bitsandbytes ships a CPU backend) but 4-bit *inference* on CPU is single-threaded and impractically slow. For CPU inference, load 4-bit, then `model.dequantize()` and cast to `torch.bfloat16`. For real use, a CUDA GPU (~6 GB VRAM) is the practical choice.
145
+ - `use_fast=False` (which the original 2024 serving code passed) is silently ignored on transformers 5.x — slow tokenizers were removed; the default fast tokenizer is correct.
146
+
147
  ## How it's used in the WIBA implementation
148
 
149
  In the WIBA serving code, this model backs the `/api/extract` endpoint at [wiba.dev](https://wiba.dev). Texts that Stage 1 classified as `Argument` are passed here to name the topic; the (text, topic) pair is then passed to Stage 3 ([stance classification](https://huggingface.co/armaniii/llama-stance-classification)) to determine whether the argument is in favor of or against that topic. For batch processing the implementation streams prompts through the pipeline with `batch_size=2` and left-padding.