kaoruhotarubi commited on
Commit
0b73ff0
·
1 Parent(s): 088e989

updated reqquirements for flash

Browse files
Files changed (1) hide show
  1. app.py +10 -19
app.py CHANGED
@@ -7,38 +7,29 @@ import os
7
  import random
8
  import re
9
  import subprocess
10
- import torch._dynamo
11
-
12
 
13
  # Define the model name
14
-
15
  OUTPUT_DIR = "output"
16
  os.makedirs(OUTPUT_DIR, exist_ok=True)
17
 
18
  model_name = "TheBloke/Amethyst-13B-Mistral-AWQ"
19
- attn_mode = "flash_attention_2" if is_flash_attn_2_available() else "default"
20
  # Load the tokenizer
21
  tokenizer = AutoTokenizer.from_pretrained(model_name)
22
 
23
- def install_flash_attn():
24
- try:
25
- import flash_attn
26
- print("✅ Flash Attention 2 is already installed.")
27
- except ImportError:
28
- print("🚀 Installing Flash Attention 2...")
29
- subprocess.run(["pip", "install", "flash-attn", "--no-build-isolation"], check=True)
30
-
31
- install_flash_attn()
32
-
33
- # Load the model
34
  model = AutoModelForCausalLM.from_pretrained(
35
  model_name,
36
- torch_dtype=torch.float16,
37
- device_map="auto",
38
- attn_implementation=attn_mode # Use Flash Attention 2 if available
39
  )
40
 
41
- model = torch.compile(model, mode="max-autotune") # Optimize inference performance
 
 
 
 
 
42
 
43
  # Define the base prompt
44
  base_prompt = """
 
7
  import random
8
  import re
9
  import subprocess
 
 
10
 
11
  # Define the model name
 
12
  OUTPUT_DIR = "output"
13
  os.makedirs(OUTPUT_DIR, exist_ok=True)
14
 
15
  model_name = "TheBloke/Amethyst-13B-Mistral-AWQ"
16
+
17
  # Load the tokenizer
18
  tokenizer = AutoTokenizer.from_pretrained(model_name)
19
 
20
+ # Load the model (without Flash Attention)
 
 
 
 
 
 
 
 
 
 
21
  model = AutoModelForCausalLM.from_pretrained(
22
  model_name,
23
+ torch_dtype=torch.float16, # Use float16 for performance
24
+ device_map="auto"
 
25
  )
26
 
27
+ # Use torch.compile() only if supported (some environments may not support it)
28
+ try:
29
+ model = torch.compile(model, mode="max-autotune") # Optimize inference performance
30
+ except RuntimeError:
31
+ print("Warning: torch.compile() is not supported on this system. Skipping optimization.")
32
+
33
 
34
  # Define the base prompt
35
  base_prompt = """