drep 4.0.2__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
drep/d_analyze.py ADDED
@@ -0,0 +1,1613 @@
1
+ #!/usr/bin/env python3
2
+ '''
3
+ d_analyze - a subset of drep
4
+
5
+ Make plots based on de-replication
6
+ '''
7
+ import matplotlib
8
+
9
+ import drep.d_cluster.cluster_utils
10
+ import drep.d_cluster.compare_utils
11
+ import drep.d_cluster.controller
12
+
13
+ matplotlib.use('Agg')
14
+
15
+ import logging
16
+ import math
17
+ import os
18
+ import sys
19
+
20
+ import pandas as pd
21
+ import seaborn as sns
22
+ import scipy.cluster.hierarchy
23
+ from sklearn import manifold
24
+ import numpy as np
25
+
26
+ from matplotlib import pyplot as plt
27
+ import matplotlib.ticker as ticker
28
+ from matplotlib.backends.backend_pdf import PdfPages
29
+ import matplotlib.patches as mpatches
30
+
31
+ import drep
32
+ import drep.d_cluster
33
+ import drep.d_filter
34
+
35
+ import traceback
36
+
37
+ import warnings
38
+ warnings.filterwarnings("ignore")#, category='all')
39
+
40
+ def d_analyze_wrapper(wd, **kwargs):
41
+ '''
42
+ Controller for the dRep analyze operation
43
+
44
+ Args:
45
+ wd: The current workDirectory
46
+ **kwargs: Command line arguments
47
+
48
+ Keyword Args:
49
+ plots: List of plots to make [list of ints, 1-6]
50
+
51
+ Returns:
52
+ Makes some plots
53
+ '''
54
+
55
+ # Load the workDirectory
56
+ wd = drep.WorkDirectory.WorkDirectory(wd)
57
+ debug = kwargs.get('debug', False)
58
+
59
+ # Figure out what plots to make
60
+ options = ['1','2','3','4','5','6']
61
+
62
+ to_plot = kwargs.get('plots', None)
63
+ to_plot = _parse_plot_options(options, to_plot)
64
+ logging.info("making plots {0}".format(', '.join(to_plot)))
65
+
66
+ # Get the plot directory
67
+ plot_dir = wd.get_dir('figures')
68
+
69
+ # 1) Primary clustering dendrogram
70
+ if '1' in to_plot:
71
+ try:
72
+ mash_dendrogram_from_wd(wd, plot_dir=plot_dir)
73
+ except BaseException as e:
74
+ logging.info('Failed to make plot #1: ' + str(e))
75
+ if debug:
76
+ traceback.print_exc()
77
+
78
+ # 2) Secondary clustering dendrogram
79
+ if '2' in to_plot:
80
+ try:
81
+ plot_secondary_dendrograms_from_wd(wd, plot_dir, **kwargs)
82
+ except BaseException as e:
83
+ logging.info('Failed to make plot #2: ' + str(e))
84
+ if debug:
85
+ traceback.print_exc()
86
+
87
+ # 3) Secondary clusters MDS
88
+ if '3' in to_plot:
89
+ try:
90
+ plot_secondary_mds_from_wd(wd, plot_dir, **kwargs)
91
+ except BaseException as e:
92
+ logging.info('Failed to make plot #3: ' + str(e))
93
+ if debug:
94
+ traceback.print_exc()
95
+
96
+ # 4) Comparison scatterplots
97
+ if '4' in to_plot:
98
+ try:
99
+ plot_scatterplots_from_wd(wd, plot_dir, **kwargs)
100
+ except BaseException as e:
101
+ logging.info('Failed to make plot #4: ' + str(e))
102
+ if debug:
103
+ traceback.print_exc()
104
+
105
+ # 5) Complex bin scorring
106
+ if '5' in to_plot:
107
+ try:
108
+ plot_binscoring_from_wd(wd, plot_dir, **kwargs)
109
+ except BaseException as e:
110
+ logging.info('Failed to make plot #5: ' + str(e))
111
+ if debug:
112
+ traceback.print_exc()
113
+
114
+ # 6) Winning plot
115
+ if '6' in to_plot:
116
+ try:
117
+ plot_winners_from_wd(wd, plot_dir, **kwargs)
118
+ except BaseException as e:
119
+ logging.info('Failed to make plot #6: ' + str(e))
120
+ if debug:
121
+ traceback.print_exc()
122
+
123
+
124
+ def mash_dendrogram_from_wd(wd, plot_dir=False):
125
+ '''
126
+ From the wd and kwargs, call plot_MASH_dendrogram
127
+
128
+ Args:
129
+ wd: WorkDirectory
130
+ plot_dir (optional): Location to store figure
131
+
132
+ Returns:
133
+ Shows plot, makes a plot in the plot_dir
134
+ '''
135
+ # Load the required data
136
+ try:
137
+ Mdb = wd.get_db('Mdb', return_none=False, forPlotting=True)
138
+ Cdb = wd.get_db('Cdb', return_none=False)
139
+ Pcluster = wd.get_primary_linkage()
140
+ Plinkage = Pcluster['linkage']
141
+ Plinkage_db = Pcluster.get('db')
142
+ clust_args = wd.arguments['cluster']
143
+ PL_thresh = clust_args.get('P_ani', False)
144
+ if PL_thresh != False:
145
+ PL_thresh = 1-PL_thresh
146
+ except:
147
+ logging.error("Skipping plot 1 - you don't have all required dataframes")
148
+ return
149
+
150
+ if 'genome_chunk' in list(Mdb.columns):
151
+ logging.error("Skipping plot 1 - cannot generate with multiround_primary_clustering enabled")
152
+ return
153
+
154
+ if Plinkage is None or isinstance(Plinkage, str):
155
+ logging.error("Skipping plot 1 - no primary linkage matrix was computed (too many genomes, or a streaming primary algorithm was used)")
156
+ return
157
+
158
+ # Leaf labels have to come from whatever the linkage was built on. The sparse
159
+ # skani Mdb only holds above-threshold pairs, so a genome with no relatives is
160
+ # absent from it and labels derived from Mdb would not match the linkage.
161
+ names = list(Plinkage_db.columns) if Plinkage_db is not None else None
162
+
163
+ # Make the plot
164
+ logging.info("Plotting primary dendrogram")
165
+ plot_MASH_dendrogram(Mdb, Cdb, Plinkage, threshold = PL_thresh,\
166
+ plot_dir = plot_dir, names = names)
167
+
168
+ def plot_secondary_dendrograms_from_wd(wd, plot_dir, **kwargs):
169
+ '''
170
+ From the wd and kwargs, make the secondary dendrograms
171
+
172
+ Args:
173
+ wd: WorkDirectory
174
+ plot_dir (optional): Location to store figure
175
+
176
+ Returns:
177
+ Makes plot
178
+ '''
179
+
180
+ # Load required databases
181
+ try:
182
+ Ndb = wd.get_db('Ndb', return_none=False)
183
+ Cdb = wd.get_db('Cdb', return_none=False)
184
+ except:
185
+ logging.error("Skipping plot 2 - you don't have all required dataframes")
186
+ return
187
+
188
+ if len(Cdb[Cdb['cluster_method'] == 'greedy']) > 0:
189
+ logging.error("Skipping plot 2 - cannot generate with greedy_secondary_clustering enabled")
190
+ return
191
+
192
+ # Initialize a .pdf
193
+ if plot_dir != False:
194
+ pp = PdfPages(plot_dir + 'Secondary_clustering_dendrograms.pdf')
195
+ save = True
196
+ else:
197
+ save = False
198
+
199
+ logging.info("Plotting secondary dendrograms")
200
+
201
+ # Load winner database if it exists
202
+ if wd.hasDb('Wdb'):
203
+ Wdb = wd.get_db('Wdb')
204
+ winners = Wdb['genome'].unique()
205
+ kwargs['winners'] = winners
206
+
207
+ # Load genome 2 taxonomy if it exists
208
+ genome2taxonomy = _get_genome2taxonomy(wd)
209
+ kwargs['genome2taxonomy'] = genome2taxonomy
210
+
211
+ # For every cluster:
212
+ for cluster in sorted(Cdb['primary_cluster'].unique()):
213
+ d = Cdb[Cdb['primary_cluster'] == cluster]
214
+
215
+ # Skip if it's a singleton
216
+ if len(d['genome'].unique()) == 1:
217
+ continue
218
+
219
+ # Load the linkage information
220
+ linkI = wd.get_cluster("secondary_linkage_cluster_{0}".format(cluster))
221
+ db = linkI['db']
222
+ linkage = linkI['linkage']
223
+ args = linkI['arguments']
224
+ threshold = args['linkage_cutoff']
225
+ alg = args['comparison_algorithm']
226
+ clust_alg = args['linkage_method']
227
+ min_cov = args['minimum_coverage']
228
+
229
+ kwargs['threshold'] = threshold
230
+ kwargs['title_string'] = 'Primary cluster {0}'.format(cluster)
231
+ kwargs['subtitle_string'] = "Comp method: {0} ".format(alg) +\
232
+ "Clust method: {0} Min cov: {1}".format(clust_alg, min_cov)
233
+
234
+ # Get name2cluster
235
+ names = list(db.columns)
236
+ name2cluster = Cdb.set_index('genome')['secondary_cluster'].to_dict()
237
+ for name in names: # Handle the case where you deleted a secondary cluster
238
+ if name not in Cdb['genome'].tolist():
239
+ name2cluster[name] = '0_0'
240
+ kwargs['name2cluster'] = name2cluster
241
+
242
+ kwargs['self_thresh'] = get_highest_self(Ndb, names)
243
+
244
+ # Make the dendrogram
245
+ _make_special_dendrogram(linkage, names, **kwargs)
246
+
247
+ # Save the file
248
+ fig = plt.gcf()
249
+ if save == True:
250
+ pp.savefig(fig)
251
+ plt.show()
252
+ plt.close(fig)
253
+
254
+ pp.close()
255
+ plt.close('all')
256
+
257
+ def plot_secondary_mds_from_wd(wd, plot_dir, **kwargs):
258
+ '''
259
+ Make a .pdf of MDS of each cluster
260
+
261
+ Args:
262
+ wd: WorkDirectory
263
+ plot_dir (optional): Location to store figure
264
+
265
+ Returns:
266
+ Makes plot
267
+ '''
268
+ # Load required databases
269
+ try:
270
+ Ndb = wd.get_db('Ndb', return_none=False)
271
+ Cdb = wd.get_db('Cdb', return_none=False)
272
+ except:
273
+ logging.error("Skipping plot 3 - you don't have all required dataframes")
274
+ return
275
+
276
+ logging.info("Plotting MDS plot")
277
+ # initialize a .pdf
278
+ if plot_dir != False:
279
+ pp = PdfPages(plot_dir + 'Secondary_clustering_MDS.pdf')
280
+ save = True
281
+ else:
282
+ save = False
283
+
284
+ # for every cluster:
285
+ Cdb = wd.get_db('Cdb')
286
+ for cluster in sorted(Cdb['primary_cluster'].unique()):
287
+ d = Cdb[Cdb['primary_cluster'] == cluster]
288
+
289
+ # Skip if it's a singleton
290
+ if len(d['genome'].unique()) == 1:
291
+ continue
292
+
293
+ # Load the linkage information
294
+ linkI = wd.get_cluster("secondary_linkage_cluster_{0}".format(cluster))
295
+ db = linkI['db']
296
+ args = linkI['arguments']
297
+
298
+ # Load name to cluster
299
+ names = list(db.columns)
300
+ name2cluster = Cdb.set_index('genome')['secondary_cluster'].to_dict()
301
+ for name in names: # Handle the case where you deleted a secondary cluster
302
+ if name not in Cdb['genome'].tolist():
303
+ name2cluster[name] = '0_0'
304
+
305
+ # Get the colors
306
+ name2color = gen_color_dictionary(names, name2cluster)
307
+ colors = [name2color[n] for n in names]
308
+
309
+ # Make cluster 2 color
310
+ cluster2color = {name2cluster[n]: name2color[n] for n in names}
311
+
312
+ # make the mds plot
313
+ _make_mds_plot(cluster, db, names, colors=colors, annotate=False,
314
+ cluster2color = cluster2color)
315
+
316
+ # save the plot
317
+ fig = plt.gcf()
318
+ if save == True:
319
+ pp.savefig(fig, bbox_inches='tight')
320
+ plt.show()
321
+ plt.close(fig)
322
+
323
+ pp.close()
324
+ plt.close('all')
325
+
326
+ def plot_scatterplots_from_wd(wd, plot_dir, **kwargs):
327
+ '''
328
+ From the wd and kwargs, call plot_scatterplots
329
+
330
+ Args:
331
+ wd: WorkDirectory
332
+ plot_dir (optional): Location to store figure
333
+
334
+ Returns:
335
+ Shows plot, makes a plot in the plot_dir
336
+ '''
337
+ # Load the required data
338
+ try:
339
+ Ndb = wd.get_db('Ndb', return_none=False)
340
+ Mdb = wd.get_db('Mdb', return_none=False)
341
+ Cdb = wd.get_db('Cdb', return_none=False)
342
+ except:
343
+ logging.error("Skipping plot 4 - you don't have all required dataframes")
344
+ return
345
+
346
+ # Make the plot
347
+ logging.info("Plotting scatterplots")
348
+ plot_scatterplots(Mdb, Ndb, Cdb, plot_dir = plot_dir)
349
+
350
+ def plot_binscoring_from_wd(wd, plot_dir, **kwargs):
351
+ '''
352
+ From the wd and kwargs, call plot_winner_scoring_complex
353
+
354
+ Args:
355
+ wd: WorkDirectory
356
+ plot_dir (optional): Location to store figure
357
+
358
+ Returns:
359
+ Shows plot, makes a plot in the plot_dir
360
+ '''
361
+ # Load the required data
362
+ try:
363
+ Sdb = wd.get_db('Sdb', return_none=False)
364
+ Cdb = wd.get_db('Cdb', return_none=False)
365
+ Wdb = wd.get_db('Wdb', return_none=False)
366
+ Bdb = wd.get_db('Bdb', return_none=False)
367
+ except:
368
+ logging.error("Skipping plot 5 - you don't have all required dataframes")
369
+ return
370
+
371
+ # Deal with genome quality
372
+ Gdb = drep.d_filter._get_run_genomeInfo(wd, Bdb, no_run=True)
373
+
374
+ # Only keep things you want
375
+ Gdb = Gdb[[c for c in Gdb.columns if c in ['genome', 'location', 'N50', 'length', 'completeness', 'contamination', 'strain_heterogeneity']]]
376
+
377
+ # Make the plot
378
+ logging.info("Plotting bin scorring plot")
379
+ plot_winner_scoring_complex(Wdb, Sdb, Cdb, Gdb, plot_dir = plot_dir, **kwargs)
380
+
381
+ def plot_winners_from_wd(wd, plot_dir, **kwargs):
382
+ '''
383
+ From the wd and kwargs, call plot_winners
384
+
385
+ Args:
386
+ wd: WorkDirectory
387
+ plot_dir: Location to store figure
388
+
389
+ Returns:
390
+ Shows plot, makes a plot in the plot_dir
391
+ '''
392
+ # Load the required data
393
+ try:
394
+ Wdb = wd.get_db('Wdb', return_none=False)
395
+ Bdb = wd.get_db('Bdb', return_none=False)
396
+ except:
397
+ logging.error("Skipping plot 6 - you don't have all required dataframes")
398
+ return
399
+
400
+ # Deal with genome quality
401
+ Gdb = drep.d_filter._get_run_genomeInfo(wd, Bdb, no_run=True)
402
+
403
+ # Only keep things you want
404
+ Gdb = Gdb[[c for c in Gdb.columns if c in ['genome', 'location', 'N50', 'length', 'completeness', 'contamination', 'strain_heterogeneity']]]
405
+
406
+ # Get optional data
407
+ Wndb = wd.get_db('Wndb')
408
+ Wmdb = wd.get_db('Wmdb')
409
+ Widb = wd.get_db('Widb')
410
+
411
+ # Make the plot
412
+ logging.info("Plotting winning genomes plot...")
413
+ plot_winners(Wdb, Gdb, Wndb, Wmdb, Widb, plot_dir = plot_dir, **kwargs)
414
+
415
+ def plot_scatterplots(Mdb, Ndb, Cdb, plot_dir=False):
416
+ '''
417
+ Make scatterplots comparing genome comparison algorithms
418
+
419
+ * plot_MASH_vs_ANIn_ani(Mdb, Ndb)
420
+ - Plot MASH_ani vs. ANIn_ani (including correlation)
421
+
422
+ * plot_MASH_vs_ANIn_cov(Mdb, Ndb)
423
+ - Plot MASH_ani vs. ANIn_cov (including correlation)
424
+
425
+ * plot_ANIn_vs_ANIn_cov(Mdb, Ndb)
426
+ - Plot ANIn vs. ANIn_cov (including correlation)
427
+
428
+ * plot_MASH_vs_len(Mdb, Ndb)
429
+ - Plot MASH_ani vs. length_difference (including correlation)
430
+
431
+ * plot_ANIn_vs_len(Ndb)
432
+ - Plot ANIn vs. length_difference (including correlation)
433
+
434
+ Args:
435
+ Mdb: DataFrame of Mash comparison results
436
+ Ndb: DataFrame of secondary clustering results
437
+ Cdb: DataFrame of Clustering results
438
+ plot_dir (optional): Location to store plot
439
+
440
+ Return:
441
+ Makes and shows plot
442
+ '''
443
+ sns.set_style('whitegrid')
444
+
445
+ # Initialize a .pdf
446
+ if plot_dir != False:
447
+ pp = PdfPages(plot_dir + 'Clustering_scatterplots.pdf')
448
+ save = True
449
+ else:
450
+ save = False
451
+
452
+ g = plot_MASH_vs_ANIn_ani(Mdb,Ndb,exclude_zero_MASH=False)
453
+ if save: pp.savefig(g)
454
+ plt.show()
455
+
456
+ g = plot_MASH_vs_secondary_ani(Mdb,Ndb,Cdb,exclude_zero_MASH=False)
457
+ if save: pp.savefig(g)
458
+ plt.show()
459
+
460
+ g = plot_MASH_vs_ANIn_cov(Mdb,Ndb)
461
+ if save: pp.savefig(g)
462
+ plt.show()
463
+
464
+ g = plot_ANIn_vs_ANIn_cov(Ndb)
465
+ if save: pp.savefig(g)
466
+ plt.show()
467
+
468
+ try:
469
+ g = plot_MASH_vs_len(Mdb,Ndb)
470
+ if save: pp.savefig(g)
471
+ plt.show()
472
+ except:
473
+ pass
474
+
475
+ try:
476
+ g = plot_ANIn_vs_len(Mdb,Ndb)
477
+ if save: pp.savefig(g)
478
+ plt.show()
479
+ except:
480
+ pass
481
+
482
+ if save:
483
+ pp.close()
484
+ plt.close('all')
485
+
486
+ def plot_MASH_vs_ANIn_ani(Mdb, Ndb, exclude_zero_MASH=True):
487
+ '''
488
+ Makes plot and retuns plt.cgf()
489
+
490
+ All parameters are obvious
491
+ '''
492
+ plt.close('all')
493
+ mdb = Mdb.copy()
494
+ mdb.rename(columns={'genome1':'querry','genome2':'reference',
495
+ 'similarity':'MASH_ANI'},inplace=True)
496
+ if exclude_zero_MASH:
497
+ mdb= mdb[mdb['MASH_ANI'] > 0]
498
+
499
+ db = pd.merge(mdb,Ndb)
500
+ db.rename(columns={'ani':'ANIn'},inplace=True)
501
+ g = sns.jointplot(x='ANIn',y='MASH_ANI',data=db)
502
+ plt.gcf().suptitle('MASH vs ANIn comparisons (all)')
503
+ plt.subplots_adjust(top=0.9)
504
+ #plt.subplots_adjust(left=0.2)
505
+ return plt.gcf()
506
+
507
+ def plot_MASH_vs_secondary_ani(Mdb,Ndb,Cdb,exclude_zero_MASH=True):
508
+ '''
509
+ Makes plot and retuns plt.cgf()
510
+
511
+ All parameters are obvious
512
+ '''
513
+ plt.close('all')
514
+ Xdb = pd.DataFrame()
515
+
516
+ # Make a database of all comparisons
517
+ mdb = Mdb.copy()
518
+ mdb.rename(columns={'genome1':'querry','genome2':'reference',
519
+ 'similarity':'MASH_ANI'},inplace=True)
520
+ if exclude_zero_MASH:
521
+ mdb= mdb[mdb['MASH_ANI'] > 0]
522
+ db = pd.merge(mdb,Ndb)
523
+ db.rename(columns={'ani':'ANIn'},inplace=True)
524
+
525
+ # Filter to only comparisons within secondary clusters
526
+ g2c = Cdb.set_index('genome')['secondary_cluster'].to_dict()
527
+ db['ref_secondary_cluster'] = db['reference'].map(g2c)
528
+ db['qu_secondary_cluster'] = db['querry'].map(g2c)
529
+ db = db[db['ref_secondary_cluster'] == db['qu_secondary_cluster']]
530
+
531
+ g = sns.jointplot(x='ANIn',y='MASH_ANI',data=db)
532
+ plt.gcf().suptitle('MASH vs ANIn comparisons (within secondary clusters only)')
533
+ plt.subplots_adjust(top=0.9)
534
+ #plt.subplots_adjust(left=0.2)
535
+ return plt.gcf()
536
+
537
+ def plot_MASH_vs_ANIn_cov(Mdb,Ndb,exclude_zero_MASH=True):
538
+ '''
539
+ Makes plot and retuns plt.cgf()
540
+
541
+ All parameters are obvious
542
+ '''
543
+ plt.close('all')
544
+ mdb = Mdb.copy()
545
+ mdb.rename(columns={'genome1':'querry','genome2':'reference',
546
+ 'similarity':'MASH_ANI'},inplace=True)
547
+ if exclude_zero_MASH:
548
+ mdb= mdb[mdb['MASH_ANI'] > 0]
549
+
550
+ db = pd.merge(mdb,Ndb)
551
+ db.rename(columns={'alignment_coverage':'ANIn_alignment_coverage'},inplace=True)
552
+ g = sns.jointplot(x='ANIn_alignment_coverage',y='MASH_ANI',data=db)
553
+ return plt.gcf()
554
+
555
+ def plot_ANIn_vs_ANIn_cov(Ndb):
556
+ '''
557
+ Makes plot and retuns plt.cgf()
558
+
559
+ All parameters are obvious
560
+ '''
561
+ plt.close('all')
562
+ db = Ndb.copy()
563
+ db.rename(columns={'alignment_coverage':'ANIn_alignment_coverage','ani':'ANIn'},inplace=True)
564
+ g = sns.jointplot(y='ANIn',x='ANIn_alignment_coverage',data=db)
565
+ return plt.gcf()
566
+
567
+ def plot_MASH_vs_len(Mdb,Ndb,exclude_zero_MASH=True):
568
+ '''
569
+ Makes plot and retuns plt.cgf()
570
+
571
+ All parameters are obvious
572
+ '''
573
+ plt.close('all')
574
+ mdb = Mdb.copy()
575
+ mdb.rename(columns={'genome1':'querry','genome2':'reference',
576
+ 'similarity':'MASH_ANI'},inplace=True)
577
+ if exclude_zero_MASH:
578
+ mdb= mdb[mdb['MASH_ANI'] > 0]
579
+
580
+ db = pd.merge(mdb,Ndb)
581
+ db['length_difference'] = abs(db['reference_length'] - db['querry_length'])
582
+ g = sns.jointplot(x='MASH_ANI',y='length_difference',data=db)
583
+
584
+ # Make a decending xaxis
585
+ axs = plt.gcf().get_axes()
586
+ plt.sca(axs[0])
587
+ plt.xlim(db['MASH_ANI'].max(),db['MASH_ANI'].min())
588
+
589
+ plt.gcf().suptitle('MASH vs length difference of genomes compared')
590
+ plt.subplots_adjust(top=0.9)
591
+ plt.subplots_adjust(left=0.2)
592
+
593
+ return plt.gcf()
594
+
595
+ def plot_ANIn_vs_len(Mdb,Ndb,exclude_zero_MASH=True):
596
+ '''
597
+ Makes plot and retuns plt.cgf()
598
+
599
+ All parameters are obvious
600
+ '''
601
+ plt.close('all')
602
+ mdb = Mdb.copy()
603
+ mdb.rename(columns={'genome1':'querry','genome2':'reference',
604
+ 'similarity':'MASH_ANI'},inplace=True)
605
+ if exclude_zero_MASH:
606
+ mdb= mdb[mdb['MASH_ANI'] > 0]
607
+
608
+ db = pd.merge(mdb,Ndb)
609
+ db['length_difference'] = abs(db['reference_length'] - db['querry_length'])
610
+ db.rename(columns={'ani':'ANIn'},inplace=True)
611
+ g = sns.jointplot(x='ANIn',y='length_difference',data=db)
612
+
613
+ # Make a decending xaxis
614
+ axs = plt.gcf().get_axes()
615
+ plt.sca(axs[0])
616
+ plt.xlim(db['ANIn'].max(),db['ANIn'].min())
617
+
618
+ plt.gcf().suptitle('ANIn vs length difference of genomes compared')
619
+ plt.subplots_adjust(top=0.9)
620
+ plt.subplots_adjust(left=0.2)
621
+ return plt.gcf()
622
+
623
+ """
624
+ CLUSETER PLOTS
625
+ """
626
+
627
+ def plot_MASH_dendrogram(Mdb, Cdb, linkage, threshold=False, plot_dir=False, names=None):
628
+ '''
629
+ Make a dendrogram of the primary clustering
630
+
631
+ Args:
632
+ Mdb: DataFrame of Mash comparison results; make sure loaded not as categories
633
+ Cdb: DataFrame of Clustering results
634
+ linkage: Result of scipy.cluster.hierarchy.linkage
635
+ threshold (optional): Line to plot on x-axis
636
+ plot_dir (optional): Location to store plot
637
+ names (optional): Leaf labels, in the order the linkage was built from.
638
+ Required when Mdb does not contain every genome -- the sparse skani
639
+ Mdb only holds above-threshold pairs, so a genome with no relatives
640
+ never appears in it and deriving labels from Mdb would silently
641
+ mismatch the linkage.
642
+
643
+ Returns:
644
+ Makes and shows plot
645
+ '''
646
+ sns.set_style('white',{'axes.grid': False})
647
+
648
+ if Mdb['genome1'].dtype.name == 'category':
649
+ logging.error("WARNING: Primary dendrogram labels may be shuffled! Load as csv to prevent this")
650
+
651
+ if names is None:
652
+ db = Mdb.pivot(index="genome1", columns="genome2", values="similarity")
653
+ names = list(db.columns)
654
+ name2cluster = Cdb.set_index('genome')['primary_cluster'].to_dict()
655
+ name2color = gen_color_dictionary(names, name2cluster)
656
+
657
+ # Make the dendrogram
658
+ g = fancy_dendrogram(linkage,names,name2color,threshold=threshold)
659
+ plt.title('MASH clustering')
660
+ plt.xlabel('MASH Average Nucleotide Identity (ANI)')
661
+
662
+ #plt.xlim([0,.4])
663
+
664
+ sns.despine(left=True,top=True,right=True,bottom=False)
665
+
666
+ # Adjust the figure size
667
+ fig = plt.gcf()
668
+ fig.set_size_inches(10,_x_fig_size(len(names),factor=.2))
669
+ plt.subplots_adjust(left=0.3)
670
+
671
+ # Adjust the x labels
672
+ plt.tick_params(axis='both', which='major', labelsize=8)
673
+ axes = plt.gca()
674
+ labels = axes.xaxis.get_majorticklocs()
675
+ for i, label in enumerate(labels):
676
+ labels[i] = float("{0:.2f}".format((1 - float(label)) * 100))
677
+ axes.set_xticklabels(labels)
678
+
679
+ # Add cluster to the y axis
680
+ g2c = Cdb.set_index('genome')['secondary_cluster'].to_dict()
681
+ axes = plt.gca()
682
+ labels = [item.get_text() for item in axes.get_yticklabels()]
683
+ for i, label in enumerate(labels):
684
+ labels[i] = "{0} ({1})".format(label, g2c[label])
685
+ axes.set_yticklabels(labels)
686
+
687
+ # Save the figure
688
+ if plot_dir != None:
689
+ plt.savefig(os.path.join(plot_dir, 'Primary_clustering_dendrogram.pdf'),\
690
+ format="pdf", transparent=True, bbox_inches='tight')
691
+ plt.show()
692
+ plt.close('all')
693
+
694
+ """
695
+ WINNER PLOTS
696
+ """
697
+
698
+ def plot_winner_scoring_complex(Wdb, Sdb, Cdb, Gdb, plot_dir= False, **kwargs):
699
+ '''
700
+ Make a plot showing the genome scoring for all genomes
701
+
702
+ Args:
703
+ Wdb: DataFrame of winning dereplicated genomes
704
+ Sdb: Scores of all genomes
705
+ Cdb: DataFrame of Clustering results
706
+ Gdb: DataFrame of genome scoring information
707
+ plot_dir (optional): Location to store plot
708
+
709
+ Returns:
710
+ makes plot
711
+ '''
712
+ # Set style
713
+ sns.reset_orig()
714
+
715
+ # Initialize a .pdf
716
+ if plot_dir != False:
717
+ pp = PdfPages(os.path.join(plot_dir, 'Cluster_scoring.pdf'))
718
+ save = True
719
+ else:
720
+ save = False
721
+
722
+ # Figure out what you're going to show
723
+ bars = _get_toshow(Gdb)
724
+ bars += ['score']
725
+
726
+ # Get winners
727
+ winners = list(Wdb['genome'].unique())
728
+
729
+ for cluster in sorted(Cdb['secondary_cluster'].unique(), key=lambda x: _comp_cluster(x)):
730
+ # Make a db for this cluster
731
+ d = Cdb[Cdb['secondary_cluster'] == cluster]
732
+ d = d.merge(Sdb, how='left', on= 'genome')
733
+ d = d.merge(Gdb, how='left', on= 'genome')
734
+ d = d[bars + ['genome']]
735
+
736
+ # Make the normalize bar plot
737
+ nd = normalize(d)
738
+ db = pd.melt(nd, id_vars=['genome'], value_vars=bars)
739
+ g = sns.barplot(data=db, y='genome', x='value', hue='variable')
740
+
741
+ # Get a list of the un-normalized values
742
+ x = pd.melt(d, id_vars=['genome'], value_vars=bars)
743
+ vals = []
744
+ for variable in x['variable'].unique():
745
+ vals += [v for v in x['value'][x['variable'] == variable].tolist()]
746
+
747
+ # # Add un-normalized values to barplots
748
+ # i = 0
749
+ # for p in g.patches:
750
+ # g.annotate("{0:.1f}".format(vals[i]), (p.get_width(), p.get_y()+(p.get_height()/1.1) ), fontsize=8)
751
+ # i += 1
752
+
753
+ plt.title('Scoring of cluster {0}'.format(cluster))
754
+ plt.xlabel('Normalized Score')
755
+ plt.legend(loc=(0,0), fancybox=True, framealpha=0.5)
756
+ plt.tick_params(axis='both', which='major', labelsize=8)
757
+
758
+ # Mark winning one
759
+ labels = d['genome'].tolist()
760
+ for i, label in enumerate(labels):
761
+ if label in winners: labels[i] = label + ' *'
762
+ axes = plt.gca()
763
+ axes.set_yticklabels(labels)
764
+
765
+ # Add taxonomy
766
+ if kwargs.get('genome2taxonomy',False) != False:
767
+ g2t = kwargs.get('genome2taxonomy')
768
+ axes = plt.gca()
769
+ labels = [item.get_text() for item in axes.get_yticklabels()]
770
+ for i, label in enumerate(labels):
771
+ labels[i] = "{0}\n{1}".format(label, g2t[label.replace(' *','')])
772
+ axes.set_yticklabels(labels)
773
+
774
+ fig = plt.gcf()
775
+ fig.set_size_inches(12,_x_fig_size(len(labels), factor=1))
776
+ plt.subplots_adjust(left=0.5)
777
+
778
+ if save == True:
779
+ pp.savefig(fig)
780
+ plt.show()
781
+ plt.close(fig)
782
+
783
+ if save:
784
+ pp.close()
785
+ plt.close('all')
786
+
787
+ def plot_winners(Wdb, Gdb, Wndb, Wmdb, Widb, plot_dir= False, **kwargs):
788
+ '''
789
+ Make a bunch of plots about the de-replicated genomes
790
+
791
+ THIS REALLY NEEDS IMPROVED UPON
792
+ '''
793
+
794
+ # Set style
795
+ sns.reset_orig()
796
+
797
+ # Initialize a .pdf
798
+ if plot_dir != False:
799
+ pp = PdfPages(plot_dir + 'Winning_genomes.pdf')
800
+ save = True
801
+ else:
802
+ save = False
803
+
804
+ # Make piecharts
805
+ if Widb is not None:
806
+ labels = []
807
+ sizes = []
808
+ for com in Widb['completeness_metric'].unique():
809
+ d = Widb[Widb['completeness_metric'] == com]
810
+ labels.append(com)
811
+ sizes.append(len(d['genome'].unique()))
812
+ labels = _annotate_labels(labels,'comp')
813
+
814
+ if (len(labels) != len(sizes)) | (len(labels) == 0): # not sure when this would happen, but it does...
815
+ logging.debug("len(labels) != len(sizes); {0} vs {1}".format(\
816
+ len(labels), len(sizes)))
817
+
818
+ else:
819
+ _make_piechart(labels,sizes)
820
+ plt.title('Overall Winner Completeness')
821
+
822
+ # Save this page
823
+ if save == True:
824
+ fig = plt.gcf()
825
+ pp.savefig(fig)
826
+ plt.show()
827
+ plt.close(fig)
828
+
829
+ labels = []
830
+ sizes = []
831
+ for com in Widb['contamination_metric'].unique():
832
+ d = Widb[Widb['contamination_metric'] == com]
833
+ labels.append(com)
834
+ sizes.append(len(d['genome'].unique()))
835
+ labels = _annotate_labels(labels,'con')
836
+ _make_piechart(labels,sizes)
837
+ plt.title('Overall Winner Contamination')
838
+
839
+ # Save this page
840
+ if save == True:
841
+ fig = plt.gcf()
842
+ pp.savefig(fig)
843
+ plt.show()
844
+ plt.close(fig)
845
+
846
+ # Figure out what you're going to show
847
+ bars = _get_toshow(Gdb)
848
+ bars += ['score']
849
+
850
+ # Make a db for the winners
851
+ d = Wdb.sort_values('score', ascending=False)
852
+ d = d.merge(Gdb, how='left', on= 'genome')
853
+ d = d[bars + ['genome']]
854
+
855
+ # Make the scoring plot
856
+ _make_scoring_plot(d,bars,**kwargs)
857
+ plt.title('Scoring of winning genomes')
858
+
859
+ # Save this page
860
+ if save == True:
861
+ fig = plt.gcf()
862
+ pp.savefig(fig)
863
+ plt.show()
864
+ plt.close(fig)
865
+
866
+ if Wmdb is not None:
867
+ # Make a MASH linkage for the dendrogram
868
+ db = Wmdb.copy()
869
+ db['dist'] = 1 - db['similarity']
870
+ linkage_db = db.pivot(index="genome1", columns="genome2", values="dist")
871
+ names = list(linkage_db.columns)
872
+ Cdb, linkage = drep.d_cluster.cluster_utils.cluster_hierarchical(linkage_db, linkage_method='average', \
873
+ linkage_cutoff= 0)
874
+
875
+ # Make the MASH dendrogram
876
+ _make_dendrogram(linkage,names)
877
+ plt.title('MASH dendrogram')
878
+
879
+ # Save this page
880
+ if save == True:
881
+ fig = plt.gcf()
882
+ pp.savefig(fig)
883
+ plt.show()
884
+ plt.close(fig)
885
+
886
+ if Wndb is not None:
887
+ # Make a ANIn linkage for the dendrogram
888
+ d = Wndb.copy()
889
+ drep.d_cluster.add_avani(d)
890
+ #d['av_ani'] = d.apply(lambda row: drep.d_cluster.average_ani (row,d),axis=1)
891
+ d['dist'] = 1 - d['av_ani']
892
+ db = d.pivot(index="reference", columns="querry", values="dist")
893
+ names = list(db.columns)
894
+ Cdb, linkage = drep.d_cluster.cluster_utils.cluster_hierarchical(db, linkage_method='average', \
895
+ linkage_cutoff= 0)
896
+
897
+ # Make the ANIn dendrogram
898
+ _make_dendrogram(linkage,names)
899
+ plt.title('ANIn dendrogram (NOT filtered for alignment length)')
900
+
901
+ # Save this page
902
+ if save == True:
903
+ fig = plt.gcf()
904
+ pp.savefig(fig)
905
+ plt.show()
906
+ plt.close(fig)
907
+
908
+ # Make a ANIn linkage for the filtered dendrogram
909
+ d = Wndb.copy()
910
+ d.loc[d['alignment_coverage'] <= 0.1, 'ani'] = 0
911
+ drep.d_cluster.add_avani(d)
912
+ #d['av_ani'] = d.apply(lambda row: drep.d_cluster.average_ani (row,d),axis=1)
913
+ d['dist'] = 1 - d['av_ani']
914
+ db = d.pivot(index="reference", columns="querry", values="dist")
915
+ names = list(db.columns)
916
+ Cdb, linkage = drep.d_cluster.cluster_utils.cluster_hierarchical(db, linkage_method='average', \
917
+ linkage_cutoff= 0)
918
+
919
+ # Make the ANIn dendrogram
920
+ _make_dendrogram(linkage,names)
921
+ plt.title('ANIn dendrogram (filtered for 10% alignment)')
922
+
923
+ # Save this page
924
+ if save == True:
925
+ fig = plt.gcf()
926
+ pp.savefig(fig)
927
+ plt.show()
928
+ plt.close(fig)
929
+
930
+ # Save the .pdf
931
+ if save:
932
+ pp.close()
933
+ plt.close('all')
934
+
935
+ def calc_dist(x1, y1, x2, y2):
936
+ '''
937
+ Return distance from two points
938
+
939
+ Args: self explainatory
940
+
941
+ Returns:
942
+ int: distance
943
+ '''
944
+ dist = math.hypot(x2 - x1, y2 - y1)
945
+ return dist
946
+
947
+ def get_highest_self(db, genomes, min = 1.0e-4):
948
+ '''
949
+ Return the highest ANI value resulting from comparing a genome to itself
950
+ '''
951
+ d = db[db['reference'].isin(genomes)]
952
+ self_thresh = 1 - d['ani'][d['reference'] == d['querry']].min()
953
+
954
+ # Because 0s don't show up on the graph
955
+ if self_thresh == float(0):
956
+ self_thresh = min
957
+ return self_thresh
958
+
959
+ def _make_piechart(labels,sizes):
960
+ '''
961
+ Used by winner plot
962
+ '''
963
+ plt.pie(sizes,labels=labels,startangle=45,\
964
+ autopct=_make_autopct(sizes),\
965
+ shadow = True)
966
+ plt.axis('equal')
967
+
968
+ def _make_autopct(values):
969
+ '''
970
+ Used by winner plot
971
+ '''
972
+ def my_autopct(pct):
973
+ total = sum(values)
974
+ val = int(round(pct*total/100.0))
975
+ return '{p:.2f}% ({v:d})'.format(p=pct,v=val)
976
+ return my_autopct
977
+
978
+ def _annotate_labels(labels,how):
979
+ '''
980
+ Used by winner plot
981
+ '''
982
+ if how == 'comp':
983
+ labs = []
984
+ for label in labels:
985
+ if label == 'near':
986
+ labs.append('near (>90%)')
987
+ if label == 'perfect':
988
+ labs.append('perfect (100%)')
989
+ if label == 'substantial':
990
+ labs.append('substantial (>70%)')
991
+ if label == 'moderate':
992
+ labs.append('moderate (>50%)')
993
+ if label == 'partial':
994
+ labs.append('partial (<50%)')
995
+ return labs
996
+
997
+ if how == 'con':
998
+ labs = []
999
+ for label in labels:
1000
+ if label == 'low':
1001
+ labs.append('low (<5%)')
1002
+ if label == 'none':
1003
+ labs.append('none (0%)')
1004
+ if label == 'medium':
1005
+ labs.append('medium (<10%)')
1006
+ if label == 'high':
1007
+ labs.append('high (<15%)')
1008
+ if label == 'very high':
1009
+ labs.append('very high (>15%)')
1010
+ return labs
1011
+
1012
+ def _x_fig_size(points, factor= .07, min= 8):
1013
+ '''
1014
+ Calculate how big the x of the figure should be
1015
+ '''
1016
+ size = points * factor
1017
+ return max([size,min])
1018
+
1019
+ def fancy_dendrogram(linkage,names,name2color=False,threshold=False,self_thresh=False):
1020
+ '''
1021
+ Make a fancy dendrogram
1022
+ '''
1023
+ # Make the dendrogram
1024
+ if threshold == False:
1025
+ scipy.cluster.hierarchy.dendrogram(linkage,labels=names,orientation='right')
1026
+ else:
1027
+ scipy.cluster.hierarchy.dendrogram(linkage,labels=names, color_threshold=threshold,\
1028
+ orientation='right')
1029
+
1030
+ # Color the names
1031
+ if name2color != False:
1032
+ ax = plt.gca()
1033
+ xlbls = ax.get_ymajorticklabels()
1034
+ for lbl in xlbls:
1035
+ color = name2color[lbl.get_text()]
1036
+ lbl.set_color('black')
1037
+ lbl.set_bbox(dict(facecolor=color, alpha=0.7, edgecolor='none', pad=2))
1038
+
1039
+ # Add the threshold
1040
+ if threshold:
1041
+ plt.axvline(x=threshold, c='k', linestyle='dotted')
1042
+ if self_thresh != False:
1043
+ plt.axvline(x=self_thresh, c='red', linestyle='dotted', lw=1)
1044
+
1045
+ g = plt.gcf()
1046
+ return g
1047
+
1048
+ def normalize(df):
1049
+ '''
1050
+ Normalize all columns in df to 0-1 except 'genome' or 'location'
1051
+
1052
+ Args:
1053
+ df: DataFrame
1054
+
1055
+ Return:
1056
+ DataFrame: Nomralized
1057
+ '''
1058
+ result = df.copy()
1059
+ for feature_name in df.columns:
1060
+ if feature_name in ['genome', 'location']:
1061
+ continue
1062
+ if not np.issubdtype(df[feature_name].dtype, np.number):
1063
+ continue
1064
+ max_value = max(df[feature_name].max(),0)
1065
+ result[feature_name] = [max((x / max_value) if max_value != 0 else 0,0) for x in result[feature_name].tolist()]
1066
+
1067
+ return result
1068
+
1069
+ def gen_color_list(names,name2cluster):
1070
+ '''
1071
+ Make a list of colors the same length as names, based on their cluster
1072
+ '''
1073
+ cm = plt.get_cmap('gist_rainbow')
1074
+
1075
+ # 1. generate cluster to color
1076
+ cluster2color = {}
1077
+ clusters = set(name2cluster.values())
1078
+ NUM_COLORS = len(clusters)
1079
+ for cluster in clusters:
1080
+ try:
1081
+ cluster2color[cluster] = cm(1.*int(cluster)/NUM_COLORS)
1082
+ except:
1083
+ cluster2color[cluster] = cm(1.*float(str(cluster).split('_')[1])/NUM_COLORS)
1084
+
1085
+ #2. generate list of colors
1086
+ colors = []
1087
+ for name in names:
1088
+ colors.append(cluster2color[name2cluster[name]])
1089
+
1090
+ return colors
1091
+
1092
+ # UC Berkeley palette. The point of coloring clusters is to tell neighbouring
1093
+ # ones apart, not to identify a cluster by its color, so a handful of distinct
1094
+ # colors cycled is strictly more readable than giving every cluster its own
1095
+ # barely-distinguishable shade.
1096
+ CLUSTER_COLORS = [
1097
+ '#003262', # Berkeley Blue
1098
+ '#FDB515', # California Gold
1099
+ '#3B7EA1', # Founders Rock
1100
+ ]
1101
+
1102
+
1103
+ def _cluster_sort_key(cluster):
1104
+ '''
1105
+ Order clusters naturally so that cycling colors lands adjacent clusters on
1106
+ different colors. Handles primary clusters ('2') and secondary clusters
1107
+ ('2_10'), sorting numerically where possible: 2_2 before 2_10, not after.
1108
+ '''
1109
+ key = []
1110
+ for part in str(cluster).split('_'):
1111
+ try:
1112
+ key.append((0, float(part), ''))
1113
+ except ValueError:
1114
+ key.append((1, 0.0, part))
1115
+ return key
1116
+
1117
+
1118
+ def gen_color_dictionary(names, name2cluster):
1119
+ '''
1120
+ Make the dictionary name2color
1121
+
1122
+ Args:
1123
+ names: key in the returned dictionary
1124
+ name2cluster: a dictionary of name to it's cluster
1125
+
1126
+ Returns:
1127
+ dict: name -> color
1128
+ '''
1129
+ # Cycle a small palette in cluster order. This is deterministic: the previous
1130
+ # implementation shuffled an unseeded colormap, so the same analysis produced
1131
+ # different colors on every run.
1132
+ clusters = sorted(set(name2cluster.values()), key=_cluster_sort_key)
1133
+ cluster2color = {c: CLUSTER_COLORS[i % len(CLUSTER_COLORS)]
1134
+ for i, c in enumerate(clusters)}
1135
+
1136
+ return {name: cluster2color[name2cluster[name]] for name in names}
1137
+
1138
+ def _comp_cluster(c):
1139
+ '''
1140
+ Used in secondary dendrogram creation
1141
+ '''
1142
+ first = int(c.split('_')[0])
1143
+ dec = float(c.split('_')[1])
1144
+ return first + dec/100
1145
+
1146
+ def _rand_cmap(nlabels, type='bright', first_color_black=True, last_color_black=False, verbose=False):
1147
+ """
1148
+ Creates a random colormap to be used together with matplotlib. Useful for segmentation tasks
1149
+ :param nlabels: Number of labels (size of colormap)
1150
+ :param type: 'bright' for strong colors, 'soft' for pastel colors
1151
+ :param first_color_black: Option to use first color as black, True or False
1152
+ :param last_color_black: Option to use last color as black, True or False
1153
+ :param verbose: Prints the number of labels and shows the colormap. True or False
1154
+ :return: colormap for matplotlib
1155
+ """
1156
+ from matplotlib.colors import LinearSegmentedColormap
1157
+ import colorsys
1158
+ import numpy as np
1159
+
1160
+
1161
+ if type not in ('bright', 'soft'):
1162
+ print ('Please choose "bright" or "soft" for type')
1163
+ return
1164
+
1165
+ if verbose:
1166
+ print('Number of labels: ' + str(nlabels))
1167
+
1168
+ # Generate color map for bright colors, based on hsv
1169
+ if type == 'bright':
1170
+ randHSVcolors = [(np.random.uniform(low=0.0, high=1),
1171
+ np.random.uniform(low=0.2, high=1),
1172
+ np.random.uniform(low=0.9, high=1)) for i in range(nlabels)]
1173
+
1174
+ # Convert HSV list to RGB
1175
+ randRGBcolors = []
1176
+ for HSVcolor in randHSVcolors:
1177
+ randRGBcolors.append(colorsys.hsv_to_rgb(HSVcolor[0], HSVcolor[1], HSVcolor[2]))
1178
+
1179
+ if first_color_black:
1180
+ randRGBcolors[0] = [0, 0, 0]
1181
+
1182
+ if last_color_black:
1183
+ randRGBcolors[-1] = [0, 0, 0]
1184
+
1185
+ random_colormap = LinearSegmentedColormap.from_list('new_map', randRGBcolors, N=nlabels)
1186
+
1187
+ # Generate soft pastel colors, by limiting the RGB spectrum
1188
+ if type == 'soft':
1189
+ low = 0.6
1190
+ high = 0.95
1191
+ randRGBcolors = [(np.random.uniform(low=low, high=high),
1192
+ np.random.uniform(low=low, high=high),
1193
+ np.random.uniform(low=low, high=high)) for i in range(nlabels)]
1194
+
1195
+ if first_color_black:
1196
+ randRGBcolors[0] = [0, 0, 0]
1197
+
1198
+ if last_color_black:
1199
+ randRGBcolors[-1] = [0, 0, 0]
1200
+ random_colormap = LinearSegmentedColormap.from_list('new_map', randRGBcolors, N=nlabels)
1201
+
1202
+ # Display colorbar
1203
+ if verbose:
1204
+ from matplotlib import colors, colorbar
1205
+ from matplotlib import pyplot as plt
1206
+ fig, ax = plt.subplots(1, 1, figsize=(15, 0.5))
1207
+
1208
+ bounds = np.linspace(0, nlabels, nlabels + 1)
1209
+ norm = colors.BoundaryNorm(bounds, nlabels)
1210
+
1211
+ cb = colorbar.ColorbarBase(ax, cmap=random_colormap, norm=norm, spacing='proportional', ticks=None,
1212
+ boundaries=bounds, format='%1i', orientation=u'horizontal')
1213
+
1214
+ return random_colormap
1215
+
1216
+
1217
+ def _parse_plot_options(options, args):
1218
+ '''
1219
+ Read user input and figure out a list of plots to make
1220
+
1221
+ Args:
1222
+ options: list of possible plots to make (default [1-6])
1223
+ args: the command line passed in
1224
+
1225
+ Returns:
1226
+ list: list of ints in the args
1227
+ '''
1228
+ to_plot = []
1229
+
1230
+ if args == []:
1231
+ return []
1232
+
1233
+ if args[0] in ['all','a']:
1234
+ to_plot += options
1235
+
1236
+ elif args == None:
1237
+ logging.error("No plots given!")
1238
+ sys.exit()
1239
+ return None
1240
+
1241
+ else:
1242
+ for arg in args:
1243
+ if arg in options:
1244
+ to_plot.append(arg)
1245
+ else:
1246
+ for letter in arg:
1247
+ if letter in options:
1248
+ to_plot.append(letter)
1249
+ else:
1250
+ logging.error("Can't interpret plotting argument {0}! quitting".format(arg))
1251
+ sys.exit()
1252
+ return to_plot
1253
+
1254
+ def _get_genome2taxonomy(wd):
1255
+ '''
1256
+ Return dictionary: genome -> taxonomy
1257
+
1258
+ All based on Bdb at the moment
1259
+
1260
+ Return False if can't do it
1261
+ '''
1262
+ try:
1263
+ Bdb = wd.get_db('Bdb')
1264
+ if 'taxonomy' in Bdb:
1265
+ genome2taxonomy = Bdb.set_index('genome')['taxonomy'].to_dict()
1266
+ return genome2taxonomy
1267
+ except:
1268
+ return False
1269
+
1270
+ def _make_dendrogram(linkage, names, **kwargs):
1271
+ '''
1272
+ Currently used by winner plot
1273
+
1274
+ names can be gotten like:
1275
+ db = db.pivot("reference","querry","ani")
1276
+ names = list(db.columns)
1277
+ '''
1278
+ threshold = kwargs.get('threshold',False)
1279
+ title = kwargs.get('title',False)
1280
+
1281
+ # Make the dendrogram
1282
+ g = fancy_dendrogram(linkage,names,threshold=threshold)
1283
+ plt.title(title)
1284
+ plt.xlabel(kwargs.get('xlabel','ANI'))
1285
+ if threshold != False:
1286
+ plt.xlim([0,2*threshold])
1287
+
1288
+ # Adjust the figure size
1289
+ fig = plt.gcf()
1290
+ fig.set_size_inches(10,_x_fig_size(len(names),factor=.2))
1291
+ plt.subplots_adjust(left=0.3)
1292
+
1293
+ # Adjust the x labels
1294
+ plt.tick_params(axis='both', which='major', labelsize=8)
1295
+ axes = plt.gca()
1296
+ labels = axes.xaxis.get_majorticklocs()
1297
+ for i, label in enumerate(labels):
1298
+ labels[i] = (1 - float(label)) * 100
1299
+ axes.set_xticklabels(labels)
1300
+
1301
+ def _make_special_dendrogram(linkage, names, **kwargs):
1302
+ '''
1303
+ Make the dendrogram used in plot 2
1304
+
1305
+ names can be gotten like:
1306
+ db = db.pivot("reference","querry","ani")
1307
+ names = list(db.columns)
1308
+
1309
+ Args:
1310
+ linkage: result of scipy.cluster.hierarchy.linkage
1311
+ names: names of the linkage
1312
+
1313
+ Kwargs:
1314
+ name2cluster: dict
1315
+ self_thresh: x-axis for soft line
1316
+ threshold: x-axis for hard line
1317
+ title_sting: title of the plot
1318
+ subtitle_string: subtitle of the plot
1319
+ winners: list of "winning" genomes (to be marked with star)
1320
+ genome2taxonomy: dictionary to add taxonomy information
1321
+
1322
+ Returns:
1323
+ Matplotlib primed with a figure
1324
+ '''
1325
+ # Load possible kwargs
1326
+ name2cluster = kwargs.get('name2cluster',False)
1327
+ self_thresh = kwargs.get('self_thresh',False)
1328
+ threshold = kwargs.get('threshold',False)
1329
+ title_string = kwargs.get('title_string','')
1330
+ subtitle_string = kwargs.get('subtitle_string','')
1331
+ winners = kwargs.get('winners',False)
1332
+ genome2taxonomy = kwargs.get('genome2taxonomy',False)
1333
+
1334
+ # Make special things
1335
+ if name2cluster != False:
1336
+ name2color = gen_color_dictionary(names, name2cluster)
1337
+ else:
1338
+ name2color = False
1339
+
1340
+ # Make the dendrogram
1341
+ sns.set_style('whitegrid')
1342
+ g = fancy_dendrogram(linkage,names,name2color,threshold=threshold,self_thresh =\
1343
+ self_thresh)
1344
+
1345
+ # Add the title and subtitle
1346
+ plt.suptitle(title_string, y=1, fontsize=18)
1347
+ plt.title(subtitle_string, fontsize=10)
1348
+
1349
+ # Adjust the x-axis
1350
+ plt.xlabel('Average Nucleotide Identity (ANI)')
1351
+ if threshold != False:
1352
+ plt.xlim([0,3*threshold])
1353
+ plt.tick_params(axis='both', which='major', labelsize=12)
1354
+ axes = plt.gca()
1355
+ labels = axes.xaxis.get_majorticklocs()
1356
+ for i, label in enumerate(labels):
1357
+ labels[i] = (1 - float(label)) * 100
1358
+ axes.set_xticklabels(labels)
1359
+ plt.gca().yaxis.grid(False)
1360
+
1361
+ # Adjust the figure size
1362
+ fig = plt.gcf()
1363
+ fig.set_size_inches(10,_x_fig_size(len(names),factor=.5))
1364
+ plt.subplots_adjust(left=0.5)
1365
+
1366
+ # Mark winning ones
1367
+ if type(winners) is not bool:
1368
+ ax = plt.gca()
1369
+ labels = [item.get_text() for item in ax.get_yticklabels()]
1370
+ for i, label in enumerate(labels):
1371
+ if label in winners: labels[i] = label + ' *'
1372
+ ax.set_yticklabels(labels)
1373
+
1374
+ # Add taxonomy
1375
+ if genome2taxonomy != False:
1376
+ g2t = genome2taxonomy
1377
+ axes = plt.gca()
1378
+ labels = [item.get_text() for item in axes.get_yticklabels()]
1379
+ for i, label in enumerate(labels):
1380
+ labels[i] = "{0}\n{1}".format(label, g2t[label.replace(' *','')])
1381
+ axes.set_yticklabels(labels)
1382
+
1383
+ def _get_toshow(Gdb):
1384
+ '''
1385
+ From Gdb, figure out what columns you can show.
1386
+ '''
1387
+ cols = list(Gdb.columns)
1388
+ cols.remove('genome')
1389
+ return cols
1390
+
1391
+ def _make_scoring_plot(db, bars,**kwargs):
1392
+ '''
1393
+ Used by winner plot
1394
+
1395
+ db is the database to plot- must contain 'genome' and all columns listed in 'bars'
1396
+ bars is all of the columns in the database to become bars
1397
+ for taxonomy, put genome2taxonomy in kwargs
1398
+ '''
1399
+ sns.reset_orig()
1400
+
1401
+ # Make the normalized bar plot
1402
+ nd = normalize(db)
1403
+ d = pd.melt(nd, id_vars=['genome'], value_vars=bars)
1404
+ g = sns.barplot(data=d, y='genome', x='value', hue='variable')
1405
+
1406
+ # Get a list of the un-normalized values
1407
+ x = pd.melt(db, id_vars=['genome'], value_vars=bars)
1408
+ vals = []
1409
+ for variable in x['variable'].unique():
1410
+ vals += [v for v in x['value'][x['variable'] == variable].tolist()]
1411
+
1412
+ # # Add un-normalized values to barplots
1413
+ # i = 0
1414
+ # for p in g.patches:
1415
+ # g.annotate("{0:.1f}".format(vals[i]), (p.get_width(), p.get_y()+(p.get_height()/1.1) ), fontsize=8)
1416
+ # i += 1
1417
+
1418
+ # Add taxonomy if available
1419
+ axes = plt.gca()
1420
+ labels = [item.get_text() for item in axes.get_yticklabels()]
1421
+ if kwargs.get('genome2taxonomy',False) != False:
1422
+ g2t = kwargs.get('genome2taxonomy')
1423
+ for i, label in enumerate(labels):
1424
+ labels[i] = "{0}\n{1}".format(label, g2t[label.replace(' *','')])
1425
+ axes.set_yticklabels(labels)
1426
+
1427
+ # Adjust labels
1428
+ plt.xlabel('Normalized Score')
1429
+ plt.legend(loc='lower right')
1430
+ plt.tick_params(axis='both', which='major', labelsize=8)
1431
+
1432
+ # Adjust figure size
1433
+ fig = plt.gcf()
1434
+ fig.set_size_inches(12,_x_fig_size(len(labels), factor=1))
1435
+ plt.subplots_adjust(left=0.5)
1436
+
1437
+ def _make_mds_plot(name, dist, names, **kwargs):
1438
+ '''
1439
+ Use MDS to cluster points.
1440
+
1441
+ Based on:
1442
+ http://baoilleach.blogspot.com/2014/01/convert-distance-matrix-to-2d.html
1443
+
1444
+ Args:
1445
+ name: title of plot
1446
+ dist: linkage databases
1447
+ names: list of names in linkage database
1448
+
1449
+ Kwargs:
1450
+ annotate: if True, write names of all points
1451
+ colors: list of colors to use
1452
+ shepard: if True, make shepard plot
1453
+ tick_spacing: default = .01
1454
+ c2c: cluster to color
1455
+
1456
+ Returns:
1457
+ Primes plot in matplotlib
1458
+ '''
1459
+
1460
+ # load kwargs
1461
+ annotate = kwargs.get('annotate', False)
1462
+ colors = kwargs.get('colors', False)
1463
+ shepard = kwargs.get('shepard', False)
1464
+ tick_spacing = kwargs.get('tick_spacing', .01)
1465
+ c2c = kwargs.get('cluster2color', False)
1466
+
1467
+ # perform MDS
1468
+ mds = manifold.MDS(n_components=2, dissimilarity="precomputed", random_state=6)
1469
+ results = mds.fit(dist)
1470
+ coords = results.embedding_
1471
+
1472
+ # make shepard plot
1473
+ if shepard:
1474
+ _shepard_plot(coords, dist, names)
1475
+
1476
+ # plot
1477
+ sns.set_style('whitegrid')
1478
+ plt.subplots_adjust(bottom = 0.1)
1479
+ plt.scatter(
1480
+ coords[:, 0], coords[:, 1], marker = 'o', linestyle = 'None',\
1481
+ color = colors
1482
+ )
1483
+
1484
+ if annotate:
1485
+ for label, x, y in zip(names, coords[:, 0], coords[:, 1]):
1486
+ plt.annotate(
1487
+ label,
1488
+ xy = (x, y), xytext = (-20, 20),
1489
+ textcoords = 'offset points', ha = 'right', va = 'bottom',
1490
+ bbox = dict(boxstyle = 'round,pad=0.5', fc = 'yellow', alpha = 0.5),
1491
+ arrowprops = dict(arrowstyle = '->', connectionstyle = 'arc3,rad=0'))
1492
+
1493
+ plt.title('Primary cluster {0}; grid spacing: {1}% ANI'.format(name, tick_spacing*100))
1494
+
1495
+ plt.axis('equal')
1496
+ ax = plt.gca()
1497
+ ax.set_xticklabels([])
1498
+ ax.set_yticklabels([])
1499
+ ax.xaxis.set_major_locator(ticker.MultipleLocator(tick_spacing))
1500
+ ax.yaxis.set_major_locator(ticker.MultipleLocator(tick_spacing))
1501
+
1502
+ # Add legend
1503
+ legendC = c2c
1504
+ if legendC != None:
1505
+ patches = [mpatches.Patch(color = color, label = label)
1506
+ for label, color in zip(legendC.keys(), legendC.values())]
1507
+ lgd = plt.legend(patches, legendC.keys(), loc = 'center left', \
1508
+ bbox_to_anchor=(1, 0.5), prop = {'size': 10, 'style': 'italic'})
1509
+
1510
+ def _shepard_plot(coords, dist, names):
1511
+ '''
1512
+ A componant of the MDS plot
1513
+ '''
1514
+ table = {'ani_dist':[], 'mds_dist':[]}
1515
+ for v1, mx, my in zip(names, coords[:, 0], coords[:, 1]):
1516
+ for v2, mx2, my2 in zip(names, coords[:, 0], coords[:, 1]):
1517
+ m_dist = calc_dist(mx, my, mx2, my2)
1518
+
1519
+ a_dist = dist.ix[v1, v2]
1520
+
1521
+ table['mds_dist'].append(m_dist)
1522
+ table['ani_dist'].append(a_dist)
1523
+
1524
+ adb = pd.DataFrame(table)
1525
+ sns.regplot(data=adb, x='ani_dist', y='mds_dist')
1526
+
1527
+
1528
+ '''
1529
+ **************************** DEPREICATED *******************************
1530
+
1531
+ * This is where that re-cluster stuff used to be
1532
+
1533
+ ################################################################################
1534
+ '''
1535
+
1536
+ def cluster_test_wrapper(wd, **kwargs):
1537
+ '''
1538
+ DEPRICATED
1539
+ '''
1540
+ # Validate arguments
1541
+ cluster = kwargs.get('cluster')
1542
+ comp_method = kwargs.get('clustering_method','ANIn')
1543
+ if comp_method not in ['ANIn','gANI']:
1544
+ raise ValueError()
1545
+ clust_method = kwargs.get('clusterAlg')
1546
+ threshold = kwargs.pop('threshold',None)
1547
+ cov_thresh = float(kwargs.get('minimum_coverage'))
1548
+ if threshold != None: threshold = 1- float(threshold)
1549
+
1550
+ # Make a bdb listing the genomes to cluster
1551
+ Cdb = wd.get_db('Cdb')
1552
+ Bdb = wd.get_db('Bdb')
1553
+ genomes = Cdb['genome'][Cdb['primary_cluster'] == int(cluster)].tolist()
1554
+ bdb = Bdb[Bdb['genome'].isin(genomes)]
1555
+
1556
+ # Get taxonomy if applicable
1557
+ if 'taxonomy' in Bdb:
1558
+ genome2taxonomy = Bdb.set_index('genome')['taxonomy'].to_dict()
1559
+ kwargs['genome2taxonomy'] = genome2taxonomy
1560
+
1561
+ # Make the comparison database
1562
+ Xdb = drep.d_cluster.compare_utils.compare_genomes(bdb, comp_method, wd, **kwargs)
1563
+
1564
+ # Remove values without enough coverage
1565
+ if comp_method == 'ANIn':
1566
+ Xdb.loc[Xdb['alignment_coverage'] <= cov_thresh, 'ani'] = 0
1567
+
1568
+ # Make it symmetrical
1569
+ drep.d_cluster.add_avani(Xdb)
1570
+ #Xdb['av_ani'] = Xdb.apply(lambda row: drep.d_cluster.average_ani (row,Xdb),axis=1)
1571
+ Xdb['dist'] = 1 - Xdb['av_ani']
1572
+ db = Xdb.pivot(index="reference", columns="querry", values="dist")
1573
+
1574
+ # Cluster it
1575
+ if threshold == None:
1576
+ threshold = float(0)
1577
+ cdb, linkage = drep.d_cluster.cluster_utils.cluster_hierarchical(db, linkage_method = clust_method, \
1578
+ linkage_cutoff = threshold)
1579
+
1580
+ # Make the plot
1581
+ names = list(db.columns)
1582
+ if comp_method == 'ANIn':
1583
+ kwargs['self_thresh'] = get_highest_self(Xdb, names)
1584
+ kwargs['threshold'] = threshold
1585
+ kwargs['title_string']="Primary_cluster_{0}_{1}".format(cluster,clust_method)
1586
+ kwargs['subtitle_string'] = "Comp method: {0} ".format(comp_method) +\
1587
+ "Clust method: {0} Min cov: {1}".format(clust_method, cov_thresh)
1588
+ kwargs['name2cluster'] = cdb.set_index('genome')['cluster'].to_dict()
1589
+
1590
+ plot_clustertest(linkage, names, wd, **kwargs)
1591
+
1592
+ def plot_clustertest(linkage, names, wd, **kwargs):
1593
+ '''
1594
+ DEPREICATED
1595
+
1596
+ names can be gotten like:
1597
+ db = db.pivot("reference","querry","ani")
1598
+ names = list(db.columns)
1599
+ '''
1600
+
1601
+ # Make the plot directory
1602
+ plot_dir = wd.location + '/figures/cluster_tests/'
1603
+ if not os.path.exists(plot_dir):
1604
+ os.makedirs(plot_dir)
1605
+
1606
+ # Make the dendrogram
1607
+ _make_special_dendrogram(linkage,names,**kwargs)
1608
+
1609
+ # Save the dendrogram
1610
+ fig = plt.gcf()
1611
+ plt.savefig("{0}{1}.pdf".format(plot_dir, kwargs['title_string']))
1612
+ plt.show()
1613
+ plt.close(fig)