JohanBeytell commited on
Commit
de71f8c
·
verified ·
1 Parent(s): 008afd4

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +21 -7
app.py CHANGED
@@ -18,7 +18,7 @@ MODEL_CONFIGS = {
18
  "Among Us": 13, "Warframe": 13, "Call of Duty": 11, "Forza Horizon": 10,
19
  "Halo": 14, "Overwatch": 9, "Subnautica": 14, "Fantasy": 16,
20
 
21
- # New Models in version 1.4
22
  "Animal Crossing": 14, "Civilization VI": 22, "Control": 22, "Cuphead": 24,
23
  "Dead Space": 18, "Diablo": 20, "Dota 2": 27, "EVE Online": 24, "GTA": 23,
24
  "Hades": 20, "Metroid": 28, "Portal": 28, "Resident Evil": 21, "RimWorld": 23,
@@ -106,17 +106,25 @@ def generate_random_name(interpreter, vocab_size, sp, max_length=10, temperature
106
  if next_token == '' or len(decoded_name) > max_length:
107
  break
108
 
 
 
 
109
  decoded_name = decoded_name.replace("▁", " ")
110
  decoded_name = decoded_name.replace("</s>", "")
111
  decoded_name = decoded_name.replace("<unk>", "")
112
- decoded_name = decoded_name.replace("<s>", "")
113
- generated_name = decoded_name.strip().capitalize()
 
 
 
 
114
 
115
- parts = generated_name.split()
 
116
  if parts and len(parts[-1]) < 3:
117
- generated_name = " ".join(parts[:-1])
118
 
119
- return generated_name.strip()
120
 
121
  def generateNames(game_type, amount, max_length=30, temperature=0.5, seed_text=""):
122
  hate_speech = detect_hate_speech(seed_text)
@@ -156,6 +164,11 @@ def generateNames(game_type, amount, max_length=30, temperature=0.5, seed_text="
156
  temperature=temperature, max_seq_len=max_seq_len
157
  )
158
  stripped = generated_name.strip()
 
 
 
 
 
159
  item_hate_speech = detect_hate_speech(stripped)
160
  item_profanity = detect_profanity([stripped], language='All')
161
  name = ''
@@ -170,7 +183,8 @@ def generateNames(game_type, amount, max_length=30, temperature=0.5, seed_text="
170
  elif item_hate_speech == ['No Hate and Offensive Speech']:
171
  name = stripped
172
 
173
- names.append(name)
 
174
 
175
  return pd.DataFrame(names, columns=['Names'])
176
 
 
18
  "Among Us": 13, "Warframe": 13, "Call of Duty": 11, "Forza Horizon": 10,
19
  "Halo": 14, "Overwatch": 9, "Subnautica": 14, "Fantasy": 16,
20
 
21
+ # New Models
22
  "Animal Crossing": 14, "Civilization VI": 22, "Control": 22, "Cuphead": 24,
23
  "Dead Space": 18, "Diablo": 20, "Dota 2": 27, "EVE Online": 24, "GTA": 23,
24
  "Hades": 20, "Metroid": 28, "Portal": 28, "Resident Evil": 21, "RimWorld": 23,
 
106
  if next_token == '' or len(decoded_name) > max_length:
107
  break
108
 
109
+ # --- TEXT NORMALIZATION BLOCK ---
110
+
111
+ # 1. Strip out unwanted tokens (including <s>)
112
  decoded_name = decoded_name.replace("▁", " ")
113
  decoded_name = decoded_name.replace("</s>", "")
114
  decoded_name = decoded_name.replace("<unk>", "")
115
+ decoded_name = decoded_name.replace("<s>", "")
116
+
117
+ # 2 & 3. Normalize spacing and apply Title Case to every word
118
+ # .split() inherently destroys all irregular spacing (multiple spaces, tabs, etc.)
119
+ words = decoded_name.split()
120
+ normalized_name = " ".join([word.capitalize() for word in words])
121
 
122
+ # 4. Split the name and check the last part length rule
123
+ parts = normalized_name.split()
124
  if parts and len(parts[-1]) < 3:
125
+ normalized_name = " ".join(parts[:-1])
126
 
127
+ return normalized_name.strip()
128
 
129
  def generateNames(game_type, amount, max_length=30, temperature=0.5, seed_text=""):
130
  hate_speech = detect_hate_speech(seed_text)
 
164
  temperature=temperature, max_seq_len=max_seq_len
165
  )
166
  stripped = generated_name.strip()
167
+
168
+ # Don't pass empty strings to validators if generation completely failed
169
+ if not stripped:
170
+ continue
171
+
172
  item_hate_speech = detect_hate_speech(stripped)
173
  item_profanity = detect_profanity([stripped], language='All')
174
  name = ''
 
183
  elif item_hate_speech == ['No Hate and Offensive Speech']:
184
  name = stripped
185
 
186
+ if name: # Only append if we actually have a name
187
+ names.append(name)
188
 
189
  return pd.DataFrame(names, columns=['Names'])
190