Spaces:
Sleeping
Sleeping
Download app.py from ddeshler/conditional-SMILES: direct link, hf CLI and curl.
- Browser
- Download file 9.47 kB
-
https://huggingface.co/spaces/ddeshler/conditional-SMILES/resolve/main/app.py
- Command line
-
hf download hf://spaces/ddeshler/conditional-SMILES/app.py
-
curl -L -o app.py https://huggingface.co/spaces/ddeshler/conditional-SMILES/resolve/main/app.py
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() | |