sunweiwei commited on
Commit
411a334
·
verified ·
1 Parent(s): eeb36f9

Upload folder using huggingface_hub

Browse files
README.md ADDED
@@ -0,0 +1,75 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # AirRep-Flan
2
+
3
+ AirRep is an attribution-friendly embedding model designed for computing training data influence on test examples.
4
+
5
+ ## Model Description
6
+
7
+ This model is based on BERT architecture (gte-small config) with an additional projection layer. It's trained to produce embeddings that can be used for:
8
+ - Text encoding
9
+ - Computing similarity scores between test and training examples
10
+ - Identifying influential training examples for test predictions
11
+
12
+ ## Model Details
13
+
14
+ - **Base Architecture**: BERT (thenlper/gte-small config)
15
+ - **Hidden Size**: 384
16
+ - **Number of Layers**: 12
17
+ - **Attention Heads**: 12
18
+ - **Max Sequence Length**: 512
19
+ - **Vocabulary Size**: 30522
20
+
21
+ ## Usage
22
+
23
+ ```python
24
+ from airrep import AirRep
25
+
26
+ # Load model
27
+ model = AirRep.from_pretrained("sunweiwei/AirRep-Flan-Small")
28
+
29
+ # Encode texts
30
+ texts = ["Question: What is AI?\nAnswer: Artificial Intelligence..."]
31
+ embeddings = model.encode(texts, batch_size=128, show_progress_bar=True)
32
+
33
+ # Compute similarity scores
34
+ test_embed = model.encode(test_texts)
35
+ train_embed = model.encode(train_texts)
36
+ scores = model.similarity(test_embed, train_embed, softmax=True)
37
+ ```
38
+
39
+ ## Installation
40
+
41
+ ```bash
42
+ pip install airrep
43
+ ```
44
+
45
+ Or install from source:
46
+
47
+ ```bash
48
+ git clone https://github.com/sunnweiwei/AirRep
49
+ cd AirRep
50
+ pip install -e .
51
+ ```
52
+
53
+ ## Training Data
54
+
55
+ This model was trained on the FLAN dataset with data influence optimization.
56
+
57
+ ## Evaluation
58
+
59
+ - **Flan LDS Spearman Correlation**: 0.21
60
+
61
+ ## Citation
62
+
63
+ If you use this model, please cite:
64
+
65
+ ```bibtex
66
+ @article{airrep2024,
67
+ title={AirRep: Attribution-friendly Representation Learning},
68
+ author={Sun, Weiwei},
69
+ year={2024}
70
+ }
71
+ ```
72
+
73
+ ## License
74
+
75
+ This model is released under the Apache 2.0 License.
config.json ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "AirRepModel"
4
+ ],
5
+ "attention_probs_dropout_prob": 0.1,
6
+ "classifier_dropout": null,
7
+ "dtype": "float32",
8
+ "hidden_act": "gelu",
9
+ "hidden_dropout_prob": 0.1,
10
+ "hidden_size": 384,
11
+ "initializer_range": 0.02,
12
+ "intermediate_size": 1536,
13
+ "layer_norm_eps": 1e-12,
14
+ "max_position_embeddings": 512,
15
+ "model_type": "airrep",
16
+ "num_attention_heads": 12,
17
+ "num_hidden_layers": 12,
18
+ "pad_token_id": 0,
19
+ "position_embedding_type": "absolute",
20
+ "transformers_version": "4.57.1",
21
+ "type_vocab_size": 2,
22
+ "use_cache": true,
23
+ "vocab_size": 30522
24
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3d6b42a8b48c6c57e549dac2d7be5692afe203b3935899aaee04e7c6453d8f2a
3
+ size 133167424
modeling_airrep.py ADDED
@@ -0,0 +1,95 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """AirRep model implementation."""
2
+
3
+ from typing import Optional
4
+ import torch
5
+ import torch.nn as nn
6
+ from transformers import BertModel, BertConfig, PreTrainedModel
7
+ from transformers.modeling_outputs import BaseModelOutput
8
+
9
+
10
+ def mean_pooling(last_hidden_states: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
11
+ """Apply mean pooling to hidden states."""
12
+ last_hidden = last_hidden_states.masked_fill(~attention_mask[..., None].bool(), 0.0)
13
+ return last_hidden.sum(dim=1) / attention_mask.sum(dim=1)[..., None]
14
+
15
+
16
+ class AirRepConfig(BertConfig):
17
+ """Configuration class for AirRep model."""
18
+
19
+ model_type = "airrep"
20
+
21
+ def __init__(
22
+ self,
23
+ **kwargs
24
+ ):
25
+ super().__init__(**kwargs)
26
+
27
+
28
+ class AirRepModel(PreTrainedModel):
29
+ """
30
+ AirRep model with BERT encoder and projection layer.
31
+
32
+ This is a standalone model, not a wrapper.
33
+ """
34
+
35
+ config_class = AirRepConfig
36
+ base_model_prefix = "airrep"
37
+
38
+ def __init__(self, config: AirRepConfig):
39
+ super().__init__(config)
40
+ self.config = config
41
+
42
+ # BERT encoder
43
+ self.bert = BertModel(config, add_pooling_layer=False)
44
+
45
+ # Projection layer
46
+ self.projector = nn.Linear(
47
+ config.hidden_size,
48
+ config.hidden_size,
49
+ dtype=torch.bfloat16
50
+ )
51
+
52
+ # Initialize weights
53
+ self.post_init()
54
+
55
+ def forward(
56
+ self,
57
+ input_ids: torch.Tensor,
58
+ attention_mask: Optional[torch.Tensor] = None,
59
+ token_type_ids: Optional[torch.Tensor] = None,
60
+ **kwargs
61
+ ) -> torch.Tensor:
62
+ """
63
+ Forward pass.
64
+
65
+ Args:
66
+ input_ids: Input token IDs
67
+ attention_mask: Attention mask
68
+ token_type_ids: Token type IDs
69
+
70
+ Returns:
71
+ Pooled and projected embeddings (batch_size, hidden_size)
72
+ """
73
+ # Get BERT outputs
74
+ outputs = self.bert(
75
+ input_ids=input_ids,
76
+ attention_mask=attention_mask,
77
+ token_type_ids=token_type_ids,
78
+ output_hidden_states=True,
79
+ return_dict=True,
80
+ )
81
+
82
+ # Mean pooling
83
+ last_hidden_state = outputs.last_hidden_state
84
+ if attention_mask is None:
85
+ attention_mask = torch.ones_like(input_ids)
86
+ pooled = mean_pooling(last_hidden_state, attention_mask)
87
+
88
+ # Project
89
+ projected = self.projector(pooled)
90
+
91
+ return projected
92
+
93
+ def save_pretrained(self, save_directory: str, **kwargs):
94
+ """Save model and config."""
95
+ super().save_pretrained(save_directory, **kwargs)
special_tokens_map.json ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "cls_token": {
3
+ "content": "[CLS]",
4
+ "lstrip": false,
5
+ "normalized": false,
6
+ "rstrip": false,
7
+ "single_word": false
8
+ },
9
+ "mask_token": {
10
+ "content": "[MASK]",
11
+ "lstrip": false,
12
+ "normalized": false,
13
+ "rstrip": false,
14
+ "single_word": false
15
+ },
16
+ "pad_token": {
17
+ "content": "[PAD]",
18
+ "lstrip": false,
19
+ "normalized": false,
20
+ "rstrip": false,
21
+ "single_word": false
22
+ },
23
+ "sep_token": {
24
+ "content": "[SEP]",
25
+ "lstrip": false,
26
+ "normalized": false,
27
+ "rstrip": false,
28
+ "single_word": false
29
+ },
30
+ "unk_token": {
31
+ "content": "[UNK]",
32
+ "lstrip": false,
33
+ "normalized": false,
34
+ "rstrip": false,
35
+ "single_word": false
36
+ }
37
+ }
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "added_tokens_decoder": {
3
+ "0": {
4
+ "content": "[PAD]",
5
+ "lstrip": false,
6
+ "normalized": false,
7
+ "rstrip": false,
8
+ "single_word": false,
9
+ "special": true
10
+ },
11
+ "100": {
12
+ "content": "[UNK]",
13
+ "lstrip": false,
14
+ "normalized": false,
15
+ "rstrip": false,
16
+ "single_word": false,
17
+ "special": true
18
+ },
19
+ "101": {
20
+ "content": "[CLS]",
21
+ "lstrip": false,
22
+ "normalized": false,
23
+ "rstrip": false,
24
+ "single_word": false,
25
+ "special": true
26
+ },
27
+ "102": {
28
+ "content": "[SEP]",
29
+ "lstrip": false,
30
+ "normalized": false,
31
+ "rstrip": false,
32
+ "single_word": false,
33
+ "special": true
34
+ },
35
+ "103": {
36
+ "content": "[MASK]",
37
+ "lstrip": false,
38
+ "normalized": false,
39
+ "rstrip": false,
40
+ "single_word": false,
41
+ "special": true
42
+ }
43
+ },
44
+ "clean_up_tokenization_spaces": true,
45
+ "cls_token": "[CLS]",
46
+ "do_basic_tokenize": true,
47
+ "do_lower_case": true,
48
+ "extra_special_tokens": {},
49
+ "mask_token": "[MASK]",
50
+ "max_length": 128,
51
+ "model_max_length": 1000000000000000019884624838656,
52
+ "never_split": null,
53
+ "pad_to_multiple_of": null,
54
+ "pad_token": "[PAD]",
55
+ "pad_token_type_id": 0,
56
+ "padding_side": "right",
57
+ "sep_token": "[SEP]",
58
+ "stride": 0,
59
+ "strip_accents": null,
60
+ "tokenize_chinese_chars": true,
61
+ "tokenizer_class": "BertTokenizer",
62
+ "truncation_side": "right",
63
+ "truncation_strategy": "longest_first",
64
+ "unk_token": "[UNK]"
65
+ }
vocab.txt ADDED
The diff for this file is too large to render. See raw diff