Commit cb52d933 authored by Bharath's avatar Bharath
Browse files

Final version of processing script

parent d5e33097
Loading
Loading
Loading
Loading
+1 −56
Original line number Diff line number Diff line
@@ -91,8 +91,7 @@ def mols_to_dict(mols, id_prefix, log_every_n=5000):
      mol_id = mol.GetProp(CID_name)
    else:
      # If mol_id is not set, then use isomeric smiles as unique identifier
      #mol_id = Chem.MolToSmiles(mol, isomericSmiles=True)
      mol_id = Chem.MolToSmiles(mol)
      mol_id = Chem.MolToSmiles(mol, isomericSmiles=True)
    mol_dict[id_prefix + mol_id] = Chem.MolToSmiles(mol, isomericSmiles=True)
  return mol_dict

@@ -178,58 +177,6 @@ def join_dict_datapoints(old_record, new_record, target_names):
      old_record[target] = new_record[target]
  return old_record

def associate_smiles_to_csv(csv_file, mol_dict, target_names):
  """Add smiles fields to targets."""
  data_df = pd.read_csv(csv_file, na_filter=False)
  data_df.fillna("")
  data_dicts = data_df.to_dict("records")
  fields = ["mol_id"] + target_names
  final_fields = ["mol_id", "smiles"] + target_names
  smiles_df = pd.DataFrame(columns=final_fields)
  for data_dict in data_dicts:
    # Trim unwanted indexing fields
    data_dict = {field: data_dict[field] for field in fields}
    mol_id = data_dict["mol_id"]
    # For now, not doing anything with num_missing
    if mol_id not in mol_dict:
      num_missing += 1
      continue
    mol_smiles = mol_dict[mol_id]
    data_dict["smiles"] = mol_smiles
    smiles_df[mol_id] = data_dict
  return smiles_df

def associate_smiles(mol_dict, csv_files, target_names, worker_pool=None):
  """Add smiles fields to all targets."""
  all_smiles_dfs = []
  if worker_pool is None:
    for csv_file in csv_files:
      smiles_df = associate_smiles_to_csv(csv_file, mol_dict, target_names)
      all_smiles_dfs.append(smiles_df)
  else:
    associate_smiles_partial = partial(
        associate_smiles_to_csv, mol_dict=mol_dict, target_names=target_names)
    all_smiles_dfs = worker_poolmap(associate_smiles_partial, csv_files)
  return all_smiles_dfs

# TODO(rbharath): This step is now the roadblock
def merge_smiles_dfs(smiles_dfs, target_names):
  """Merge data from target and molecule listings."""
  #print("len(mol_dict) = %d" % len(mol_dict))
  merged_data = {}
  merge_pos, merge_map = 0, {}
  for ind, smiles_df in enumerate(smiles_dfs):
    print("Merging %d/%d targets" % (ind, len(smiles_dfs)))
    data_dicts = smiles_df.to_dict("records")
    for data_dict in data_dicts:
      mol_id = data_dict["mol_id"]
      if mol_id not in merged_data:
        merged_data[mol_id] = data_dict
      else:
        merged_data[mol_id] = join_dict_datapoints(
            merged_data[mol_id], data_dict, target_names)
  return merged_data

def merge_mol_data_dicts(mol_dict, csv_files, target_names):
  """Merge data from target and molecule listings."""
  print("len(mol_dict) = %d" % len(mol_dict))
@@ -280,8 +227,6 @@ def generate_csv(data_dir, id_prefix, out, overwrite, worker_pool=None):

  csv_files = process_targets(targets_dir, overwrite, worker_pool)

  #smiles_dfs = associate_smiles(mol_dict, csv_files, target_names, worker_pool)
  #merged_dict = merge_smiles_dfs(smiles_dfs, target_names)
  merged_dict = merge_mol_data_dicts(mol_dict, csv_files, target_names)
  merged_df = pd.DataFrame(merged_dict.values())
  merged_df.fillna("")