GenerTeam commited on
Commit
14a1105
·
verified ·
1 Parent(s): c0bb547

Update README.md

Browse files
Files changed (1) hide show
  1. README.md +44 -66
README.md CHANGED
@@ -27,119 +27,100 @@ For more technical details, please refer to our paper [GENERator: A Long-Context
27
 
28
 
29
  ## How to use
30
- ### Simple example1: generation
31
 
32
  ```python
33
 
34
  import torch
35
  from transformers import AutoTokenizer, AutoModelForCausalLM
36
 
37
- # Load the tokenizer and model.
38
- tokenizer = AutoTokenizer.from_pretrained("GenerTeam/GENERator-eukaryote-3b-base", trust_remote_code=True)
39
- model = AutoModelForCausalLM.from_pretrained("GenerTeam/GENERator-eukaryote-3b-base")
40
- config = model.config
 
 
41
 
42
- max_length = config.max_position_embeddings
 
 
 
43
 
44
  # Define input sequences.
45
  sequences = [
46
- "ATGAGGTGGCAAGAAATGGGCTAC",
47
- "GAATTCCATGAGGCTATAGAATAATCTAAGAGAAAT"
48
  ]
49
 
50
- def left_padding(sequence, padding_char='A', multiple=6):
51
- remainder = len(sequence) % multiple
52
- if remainder != 0:
53
- padding_length = multiple - remainder
54
- return padding_char * padding_length + sequence
55
- return sequence
56
-
57
- def left_truncation(sequence, multiple=6):
58
- remainder = len(sequence) % multiple
59
- if remainder != 0:
60
- return sequence[remainder:]
61
- return sequence
62
-
63
- # Apply left_padding to all sequences
64
- # padded_sequences = [left_padding(seq) for seq in sequences]
65
-
66
- # Apply left_truncation to all sequences
67
- truncated_sequences = [left_truncation(seq) for seq in sequences]
68
-
69
- # Process the sequences
70
- sequences = [tokenizer.bos_token + sequence for sequence in truncated_sequences]
71
 
72
  # Tokenize the sequences
73
- tokenizer.padding_side = "left"
74
  inputs = tokenizer(
75
- sequences,
76
  add_special_tokens=False,
77
  return_tensors="pt",
78
  padding=True,
79
- truncation=True,
80
- max_length=max_length
81
- )
82
 
83
  # Generate the sequences
84
  with torch.inference_mode():
85
- outputs = model.generate(**inputs, max_new_tokens=32, temperature=0.00001, top_k=1)
86
 
87
  # Decode the generated sequences
88
  decoded_sequences = tokenizer.batch_decode(outputs, skip_special_tokens=True)
89
 
90
  # Print the decoded sequences
91
  print(decoded_sequences)
92
-
93
- # It is expected to observe non-sense decoded sequences (e.g., 'AAAAAA')
94
- # The input sequences are too short to provide sufficient context.
95
  ```
96
 
97
- ### Simple example2: embedding
98
 
99
  ```python
100
 
101
-
102
  import torch
103
  from transformers import AutoTokenizer, AutoModelForCausalLM
104
 
105
- # Load the tokenizer and model
106
- tokenizer = AutoTokenizer.from_pretrained("GENERator-eukaryote-3b-base", trust_remote_code=True)
107
- model = AutoModelForCausalLM.from_pretrained("GENERator-eukaryote-3b-base")
 
 
 
108
 
109
- # Get model configuration
110
- config = model.config
111
- max_length = config.max_position_embeddings
 
112
 
113
- # Define input sequences
114
  sequences = [
115
- "ATGAGGTGGCAAGAAATGGGCTAC",
116
- "GAATTCCATGAGGCTATAGAATAATCTAAGAGAAAT"
117
  ]
118
 
119
  # Truncate each sequence to the nearest multiple of 6
120
- processed_sequences = [tokenizer.bos_token + seq[:len(seq)//6*6] for seq in sequences]
121
 
122
- # Tokenization
123
- tokenizer.padding_side = "right"
124
  inputs = tokenizer(
125
  processed_sequences,
126
- add_special_tokens=True,
127
  return_tensors="pt",
128
  padding=True,
129
- truncation=True,
130
- max_length=max_length
131
- )
132
 
133
- # Model Inference
134
  with torch.inference_mode():
135
  outputs = model(**inputs, output_hidden_states=True)
136
 
137
  hidden_states = outputs.hidden_states[-1]
138
  attention_mask = inputs["attention_mask"]
139
 
140
- # Option 1: Last token (EOS) embedding
141
  last_token_indices = attention_mask.sum(dim=1) - 1
142
- eos_embeddings = hidden_states[torch.arange(hidden_states.size(0)), last_token_indices, :]
143
 
144
  # Option 2: Mean pooling over all tokens
145
  expanded_mask = attention_mask.unsqueeze(-1).expand(hidden_states.size()).to(torch.float32)
@@ -147,16 +128,13 @@ sum_embeddings = torch.sum(hidden_states * expanded_mask, dim=1)
147
  mean_embeddings = sum_embeddings / expanded_mask.sum(dim=1)
148
 
149
  # Output
150
- print("EOS (Last Token) Embeddings:", eos_embeddings)
151
  print("Mean Pooling Embeddings:", mean_embeddings)
152
 
153
  # ============================================================================
154
- # Additional notes:
155
- # - The preprocessing step ensures sequences are multiples of 6 for 6-mer tokenizer
156
- # - For causal LM, the last token embedding (EOS) is commonly used
157
- # - Mean pooling considers all tokens including BOS and content tokens
158
- # - The choice depends on your downstream task requirements
159
- # - Both methods handle variable sequence lengths via attention mask
160
  # ============================================================================
161
 
162
  ```
 
27
 
28
 
29
  ## How to use
30
+ ### Example 1: Sequence Generation
31
 
32
  ```python
33
 
34
  import torch
35
  from transformers import AutoTokenizer, AutoModelForCausalLM
36
 
37
+ model = AutoModelForCausalLM.from_pretrained(
38
+ "GenerTeam/GENERator-eukaryote-1.2b-base",
39
+ attn_implementation="flash_attention_2",
40
+ trust_remote_code=True,
41
+ dtype=torch.bfloat16,
42
+ ).cuda().eval()
43
 
44
+ tokenizer = AutoTokenizer.from_pretrained(
45
+ "GenerTeam/GENERator-eukaryote-1.2b-base",
46
+ trust_remote_code=True,
47
+ )
48
 
49
  # Define input sequences.
50
  sequences = [
51
+ "ATCGATCGATCGATCGATCGATCGATCGATCGATCGATCGATCGATCGATCGATCGATCG",
52
+ "ACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGT"
53
  ]
54
 
55
+ # Truncate each sequence to the nearest multiple of 6
56
+ processed_sequences = ["<s>" + seq[len(seq)%6:] for seq in sequences]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
57
 
58
  # Tokenize the sequences
 
59
  inputs = tokenizer(
60
+ processed_sequences,
61
  add_special_tokens=False,
62
  return_tensors="pt",
63
  padding=True,
64
+ padding_side="left",
65
+ ).to("cuda")
 
66
 
67
  # Generate the sequences
68
  with torch.inference_mode():
69
+ outputs = model.generate(**inputs, max_new_tokens=32, do_sample=False)
70
 
71
  # Decode the generated sequences
72
  decoded_sequences = tokenizer.batch_decode(outputs, skip_special_tokens=True)
73
 
74
  # Print the decoded sequences
75
  print(decoded_sequences)
 
 
 
76
  ```
77
 
78
+ ### Example 2: Embedding Extraction
79
 
80
  ```python
81
 
 
82
  import torch
83
  from transformers import AutoTokenizer, AutoModelForCausalLM
84
 
85
+ model = AutoModelForCausalLM.from_pretrained(
86
+ "GenerTeam/GENERator-eukaryote-1.2b-base",
87
+ attn_implementation="flash_attention_2",
88
+ trust_remote_code=True,
89
+ dtype=torch.bfloat16,
90
+ ).cuda().eval()
91
 
92
+ tokenizer = AutoTokenizer.from_pretrained(
93
+ "GenerTeam/GENERator-eukaryote-1.2b-base",
94
+ trust_remote_code=True,
95
+ )
96
 
97
+ # Define input sequences.
98
  sequences = [
99
+ "ATCGATCGATCGATCGATCGATCGATCGATCGATCGATCGATCGATCGATCGATCGATCG",
100
+ "ACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGT"
101
  ]
102
 
103
  # Truncate each sequence to the nearest multiple of 6
104
+ processed_sequences = ["<s>" + seq[len(seq)%6:] for seq in sequences]
105
 
106
+ # Tokenize the sequences
 
107
  inputs = tokenizer(
108
  processed_sequences,
109
+ add_special_tokens=False,
110
  return_tensors="pt",
111
  padding=True,
112
+ padding_side="right",
113
+ ).to("cuda")
 
114
 
 
115
  with torch.inference_mode():
116
  outputs = model(**inputs, output_hidden_states=True)
117
 
118
  hidden_states = outputs.hidden_states[-1]
119
  attention_mask = inputs["attention_mask"]
120
 
121
+ # Option 1: Last token embedding
122
  last_token_indices = attention_mask.sum(dim=1) - 1
123
+ last_token_embeddings = hidden_states[torch.arange(hidden_states.size(0)), last_token_indices, :]
124
 
125
  # Option 2: Mean pooling over all tokens
126
  expanded_mask = attention_mask.unsqueeze(-1).expand(hidden_states.size()).to(torch.float32)
 
128
  mean_embeddings = sum_embeddings / expanded_mask.sum(dim=1)
129
 
130
  # Output
131
+ print("Last Token Embeddings:", last_token_embeddings)
132
  print("Mean Pooling Embeddings:", mean_embeddings)
133
 
134
  # ============================================================================
135
+ # The choice depends on your downstream task requirements
136
+ # - Last token embeddings capture more localized gene-level information (e.g., strand, codon phase).
137
+ # - Mean pooling embeddings capture species-level information.
 
 
 
138
  # ============================================================================
139
 
140
  ```