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()