Dylan Deshler
conditional generation
fcd6227
Raw History Blame Contribute Delete
9.47 kB
import gradio as gr
import torch
import numpy as np
import gradio as gr
from rdkit import Chem
from rdkit.Chem import Draw, Descriptors, rdMolDescriptors, QED
from PIL import Image, ImageDraw, ImageFont
from llama_model import ConditionalTransformer, ModelArgs
from generate_no_cache import generate as conditional_generate
from transformers import PreTrainedTokenizerFast
def decode(out, is_cond=False):
'''Decodes and returns generations form a batch '''
out = [tokenizer.decode(o) for o in out]
out = [''.join(o.split(' ')) for o in out]
if is_cond:
out = [o.split('[SEP]')[0] for o in out]
return out
else:
outs = []
for o in out:
outs.extend(o.split('[SEP]')[1:-1])
np.random.shuffle(outs)
return outs
def verify_smiles(smiles: str) -> bool:
"""Returns True if the SMILES is valid, False otherwise."""
try:
mol = Chem.MolFromSmiles(smiles)
return mol is not None
except:
return False
def generate_mol_background_color(mol) -> tuple:
logp = Descriptors.MolLogP(mol)
h_donors = Descriptors.NumHDonors(mol)
mol_wt = Descriptors.MolWt(mol)
red = min(int(logp / 6 * 255), 255) # logP: -3 to +3 → 0–255
green = min(h_donors * 30, 255) # H-donors: 0–8
blue = min(int(mol_wt / 500 * 255), 255) # MW: ~0–500
return (red, green, blue)
def generate_baseball_card(smiles: str, img_size=(300, 300), font_path=None) -> Image.Image:
mol = Chem.MolFromSmiles(smiles)
if mol is None:
raise ValueError("Invalid SMILES")
# Get molecule-specific background color
bg_color = generate_mol_background_color(mol)
# Create transparent molecule image
mol_img = Draw.MolToImage(mol, size=img_size, kekulize=True, bgColor=(255, 255, 255))
# Property summary
props = {
"Formula": rdMolDescriptors.CalcMolFormula(mol),
"Mol Weight": f"{Descriptors.MolWt(mol):.2f}",
"LogP": f"{Descriptors.MolLogP(mol):.2f}",
"H-Donors": Descriptors.NumHDonors(mol),
"H-Acceptors": Descriptors.NumHAcceptors(mol),
"Rotatable Bonds": Descriptors.NumRotatableBonds(mol),
"QED": round(QED.qed(mol), 2),
}
# Font
font_size = 20
font = ImageFont.truetype(font_path, font_size) if font_path else ImageFont.load_default()
# Layout dimensions
padding = 20
box_margin = 6
stats_height = len(props) * (font_size + 6) + padding
card_width = img_size[0] + 2 * padding
card_height = img_size[1] + stats_height + 3 * padding
# Create background card
card = Image.new("RGB", (card_width, card_height), color=bg_color)
draw = ImageDraw.Draw(card)
# --- Draw Molecule Image Box ---
img_x = padding
img_y = padding
mol_box_coords = [
img_x - box_margin,
img_y - box_margin,
img_x + img_size[0] + box_margin - 1,
img_y + img_size[1] + box_margin - 1
]
draw.rectangle(mol_box_coords, fill="white", outline="black", width=2)
card.paste(mol_img, (img_x, img_y))
# line_width = draw.textlength(smiles, font=font)
# draw.text(((card_width - line_width) // 2, img_y), smiles, fill="black", font=font)
# --- Draw Stats Box ---
stats_y = img_y + img_size[1] + padding
stats_box_coords = [
padding - box_margin,
stats_y - box_margin,
card_width - padding + box_margin - 1,
stats_y + stats_height - padding + box_margin - 1
]
draw.rectangle(stats_box_coords, fill="white", outline="black", width=2)
# Write centered stats
text_y = stats_y + 6
for key, val in props.items():
line = f"{key}: {val}"
line_width = draw.textlength(line, font=font)
draw.text(((card_width - line_width) // 2, text_y), line, fill="black", font=font)
text_y += font_size + 6
return card
def generate_default_card(img_size=(300, 300)):
font_size = 20
padding = 20
n_stats = 7
stats_height = n_stats * (font_size + 6) + padding
card_width = img_size[0] + 2 * padding
card_height = img_size[1] + stats_height + 3 * padding
card = Image.new("RGB", (card_width, card_height), color='gray')
return card
# Load model config and checkpoint
tokenizer = PreTrainedTokenizerFast(tokenizer_file='tokenizer_4096.json')
checkpoint = torch.load('ckpt.pt', map_location=torch.device('cpu'))
qed_bins = np.load('20_bins.npy') # maps [0, 1] -> [1, 20] by pdf and idx 0 is reserved for null embedding
args = ModelArgs(**checkpoint['model_args'])
model = ConditionalTransformer(args)
model.load_state_dict(checkpoint['model'])
model.eval()
# O=C(C)Oc1ccccc1C(=O)O
def generate_continuations(text):
idx = [4096, 3] + tokenizer.encode(text)
idx = torch.from_numpy(np.asarray(idx)).long().unsqueeze(0).expand(4, -1)
out = model.generate(idx=idx, max_new_tokens=64, temperature=1.0).cpu().detach().numpy()
# only the first generation in eac batch index is a continuation
out = [tokenizer.decode(o) for o in out]
out = [''.join(o.split(' ')) for o in out]
out = [o.split('[SEP]')[1] for o in out]
outs = []
for o in out:
if verify_smiles(o):
outs.append(generate_baseball_card(o))
if len(outs) < 3:
outs.extend([generate_default_card() for _ in range(3 - len(outs))])
return outs[:3]
def generate_similar(text):
idx = [4096, 3] + tokenizer.encode(text)
idx = torch.from_numpy(np.asarray(idx)).long().unsqueeze(0).expand(4, -1)
out = model.generate(idx=idx, max_new_tokens=64, temperature=1.0).cpu().detach().numpy()
out = decode(out, is_cond=False)
outs = []
for o in out:
if verify_smiles(o):
outs.append(generate_baseball_card(o))
if len(outs) < 3:
outs.extend([generate_default_card() for _ in range(3 - len(outs))])
return outs[:3]
cond_map = {
'Low': 0,
'Medium': 0.5,
'High': 1
}
def generate_cfg(qed):
if qed == 'None' or qed is None:
sep = torch.from_numpy(np.asarray([4096, 3])).long().unsqueeze(0).expand(2, -1)
out = model.generate(idx=sep, max_new_tokens=64, temperature=1.0).cpu().detach().numpy()
out = decode(out, is_cond=False)
elif qed in cond_map:
sep = torch.from_numpy(np.asarray([4096 + np.digitize(cond_map[qed], qed_bins), 3])).long().unsqueeze(0).expand(4, -1)
out = conditional_generate(model, sep, max_new_tokens=64, cfg_scale=3).cpu().detach().numpy()
out = decode(out[:, 2:], is_cond=True)
else:
raise NotImplementedError(f'QED = {qed} not implemented in generate')
outs = []
for o in out:
if verify_smiles(o):
outs.append(generate_baseball_card(o))
if len(outs) < 3:
outs.extend([generate_default_card() for _ in range(3 - len(outs))])
return outs[:3]
def generate(qed, text, text_type):
# figure out how to generate
if qed == 'None' or qed is None:
if text != '':
if text_type == 'Continuation':
return generate_continuations(text)
elif text_type == 'Similar':
return generate_similar(text)
elif text_type == 'None' or text_type is None:
sep = torch.from_numpy(np.asarray([4096, 3])).long().unsqueeze(0).expand(2, -1)
out = model.generate(idx=sep, max_new_tokens=64, temperature=1.0).cpu().detach().numpy()
out = decode(out, is_cond=False)
else:
raise NotImplementedError(f'Text Type {text_type} is not implemented in generate')
else:
sep = torch.from_numpy(np.asarray([4096, 3])).long().unsqueeze(0).expand(2, -1)
out = model.generate(idx=sep, max_new_tokens=64, temperature=1.0).cpu().detach().numpy()
out = decode(out, is_cond=False)
elif qed in cond_map:
sep = torch.from_numpy(np.asarray([4096 + np.digitize(cond_map[qed], qed_bins), 3])).long().unsqueeze(0).expand(4, -1)
out = conditional_generate(model, sep, max_new_tokens=64, cfg_scale=3).cpu().detach().numpy()
out = decode(out[:, 2:], is_cond=True)
else:
raise NotImplementedError(f'QED = {qed} not implemented in generate')
outs = []
for o in out:
if verify_smiles(o):
outs.append(generate_baseball_card(o))
if len(outs) < 3:
outs.extend([generate_default_card() for _ in range(3 - len(outs))])
return outs[:3]
with gr.Blocks() as demo:
gr.Markdown("# A Lightweight, Conditional, and Steerable Molecule Generator")
with gr.Row():
btn = gr.Button("Generate Molecules")
with gr.Row():
img1 = gr.Image()
img2 = gr.Image()
img3 = gr.Image()
with gr.Accordion("Set optional conditions for molecular generation. For now, only one conditional option is allowed at a time!", open=True):
qed = gr.Radio(['None', 'Low', 'Medium', 'High'], value='None', label="Quantitative Estimate of Druglikeness (QED)", show_label=True)
text = gr.Textbox('O=C(C)Oc1ccccc1C', label='Type SMILES to condition generations on', show_label=True)
text_button = gr.Radio(['None', 'Continuation', 'Similar'], value='None', label='How should the above SMILES be used?', show_label=True)
btn.click(
fn=generate,
inputs=[qed, text, text_button],
outputs=[img1, img2, img3]
)
demo.launch()